diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 92230fc8892..7fd66e3325e 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -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' diff --git a/.github/workflows/mutation-test.yml b/.github/workflows/mutation-test.yml index 6684952b998..da4fe073a6a 100644 --- a/.github/workflows/mutation-test.yml +++ b/.github/workflows/mutation-test.yml @@ -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: diff --git a/Dockerfile b/Dockerfile index 9977ebb82d7..a127cdabd59 100644 --- a/Dockerfile +++ b/Dockerfile @@ -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 \ diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 34c9c606991..9ee076ce825 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -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 \ diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index 8e05f312ba0..946b4de6f5e 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -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 diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index d96f3e110d9..89a44fcdeef 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -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: diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index cf9dafcb222..57b05c9bec8 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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) diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index fea55cd1db4..12465377b51 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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 diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index c93f95ec97d..fa41070a8de 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -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 ), diff --git a/litellm/integrations/otel/plumbing/context.py b/litellm/integrations/otel/plumbing/context.py index 8acac112c3d..939559347b1 100644 --- a/litellm/integrations/otel/plumbing/context.py +++ b/litellm/integrations/otel/plumbing/context.py @@ -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: diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 88dddb59cc7..cecc35ee1c1 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -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 diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 7000c20d9c4..90f735707bf 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -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, ) diff --git a/litellm/llms/vertex_ai/batches/transformation.py b/litellm/llms/vertex_ai/batches/transformation.py index 6bbe8f75701..df903ba7ef0 100644 --- a/litellm/llms/vertex_ai/batches/transformation.py +++ b/litellm/llms/vertex_ai/batches/transformation.py @@ -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": diff --git a/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py b/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py index 7122c64ec64..f7bc14575c7 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py +++ b/litellm/proxy/_experimental/mcp_server/auth/litellm_auth_handler.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4fca4406a6f..483f57f9139 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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( diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index a9c2a12aff7..33bca782e0b 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -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 = {} diff --git a/litellm/proxy/config_resolvers/sso.py b/litellm/proxy/config_resolvers/sso.py index 3d83c06dd62..97c42106018 100644 --- a/litellm/proxy/config_resolvers/sso.py +++ b/litellm/proxy/config_resolvers/sso.py @@ -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"), ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 2d67c22f0aa..ca3bb0ee361 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 28b9dec100f..f2c5a95202b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -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" diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py index 5c9f93fc2cd..717f5b6c5fe 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py +++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py @@ -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)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 8e3abfbf159..d4d23cd2e37 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 1f2c9e0c182..bd00e9815a8 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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"), diff --git a/litellm/proxy/management_endpoints/sso/saml_sso.py b/litellm/proxy/management_endpoints/sso/saml_sso.py new file mode 100644 index 00000000000..37b641ca123 --- /dev/null +++ b/litellm/proxy/management_endpoints/sso/saml_sso.py @@ -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 [] diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 31b98bf20e4..8682b61f910 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index acb2e50c79b..9364d7eae3a 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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)) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 25fa0819930..1f5aacc2115 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5b81d1f2da3..171e17ce650 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 12c890ec91d..429ddeef36a 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -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], diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index c86794b90f8..a324e71e289 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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", diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 9f689a2dd31..314bb653196 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -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]] diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py index b4375237917..0e816985cb0 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py @@ -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 diff --git a/litellm/types/proxy/management_endpoints/ui_sso.py b/litellm/types/proxy/management_endpoints/ui_sso.py index 742e0f7818f..d4b1d98f957 100644 --- a/litellm/types/proxy/management_endpoints/ui_sso.py +++ b/litellm/types/proxy/management_endpoints/ui_sso.py @@ -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, diff --git a/pyproject.toml b/pyproject.toml index a448ab042b8..44c1967ad9b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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'", diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index d3d70ff5ff4..39e9bf4773d 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -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 diff --git a/tests/e2e/coverage_registry/mcp.yaml b/tests/e2e/coverage_registry/mcp.yaml index d477b257cb0..ab644118a47 100644 --- a/tests/e2e/coverage_registry/mcp.yaml +++ b/tests/e2e/coverage_registry/mcp.yaml @@ -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 diff --git a/tests/e2e/mcp/datadog_mcp.py b/tests/e2e/mcp/datadog_mcp.py index f1b9461f23b..d1ea53a0b3b 100644 --- a/tests/e2e/mcp/datadog_mcp.py +++ b/tests/e2e/mcp/datadog_mcp.py @@ -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 diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index b0aa4c68e3a..f758a41cae6 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -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( diff --git a/tests/e2e/mcp/test_mcp_access_group_e2e.py b/tests/e2e/mcp/test_mcp_access_group_e2e.py new file mode 100644 index 00000000000..d7ff9736896 --- /dev/null +++ b/tests/e2e/mcp/test_mcp_access_group_e2e.py @@ -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}" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index d21920cf848..af695acaa5e 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -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): diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index ee18c96c393..d2218b08386 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -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 diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index de5b3b32105..a2e18a62638 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -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 diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 2789b4e61d5..a111b932f2c 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -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 diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 5f6002f4cdf..d3593f2c06b 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -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 ] diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 64813c1eda7..bac0ae54033 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -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) == [] diff --git a/tests/test_litellm/integrations/test_guardrail_logging_sync.py b/tests/test_litellm/integrations/test_guardrail_logging_sync.py index 5dcd1114b3d..f9e1a3efbd0 100644 --- a/tests/test_litellm/integrations/test_guardrail_logging_sync.py +++ b/tests/test_litellm/integrations/test_guardrail_logging_sync.py @@ -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(): diff --git a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py index ace9399cf53..c3e9d67ddad 100644 --- a/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py +++ b/tests/test_litellm/integrations/test_otel_guardrail_violation_spans.py @@ -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), diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index b67ea91bb0b..b4f539da286 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -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.""" diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 9cd1fbb59a6..48acdd348e9 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -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""" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index 3327fc39f73..8875a75e86f 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -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. diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 7bd1d7a6031..87d67e0e8b7 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -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", [ diff --git a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py index f8bd83fc7df..1043c26c6ec 100644 --- a/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py +++ b/tests/test_litellm/llms/pass_through/guardrail_translation/test_handler.py @@ -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 - ) diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py index 71da1d39876..1b37ade6b30 100644 --- a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py @@ -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({}) == "" diff --git a/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py b/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py index 20bea98351f..9f91d9ee2c9 100644 --- a/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py +++ b/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py @@ -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. diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py index fe6cb98d1f5..9002d1f81a3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py @@ -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"] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py index 0358ca998aa..914af0e2368 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_openai_moderation_streaming.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index 7f412c008ca..776df985d46 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -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(): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 18b5bd92411..89b6af27719 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -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.""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py index ca57118ee9d..36a2e205ea7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e84e9b74201..bf904dbe394 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 26feddadf79..14cab50f441 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index 83593c20110..71e775842e3 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_saml_sso.py b/tests/test_litellm/proxy/management_endpoints/test_saml_sso.py new file mode 100644 index 00000000000..57decc7d458 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_saml_sso.py @@ -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 ( + '' + f'' + '' + '' + f"{cert_body}" + "" + '' + "" + ) + + +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'' + + "".join(f"{v}" for v in values) + + "" + for name, values in attributes.items() + ) + + assertion = ( + '' + f"{IDP_ENTITY}" + "" + '' + f"{email}" + '' + f'' + "" + f'' + f"{SP_ENTITY}" + "" + f'' + "" + "urn:oasis:names:tc:SAML:2.0:ac:classes:Password" + "" + f"{attr_xml}" + "" + ) + + if sign: + signed = OneLogin_Saml2_Utils.add_sign(assertion, key_pem, cert_pem) + assertion = (signed.decode() if isinstance(signed, bytes) else signed).replace( + '', "" + ) + + return ( + '' + '' + f"{IDP_ENTITY}" + '' + "" + f"{assertion}" + ) + + +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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 63a47428780..795b7cd5a9e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index bad76864ca7..087aaec9215 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index f506b9665a6..93b3ef1cce8 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -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() diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 85dbf70b452..20451f5d0ac 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -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 ): diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 64c14abfd83..5c711fc6c34 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -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 diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py index 6a339b37a80..715d66db181 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py @@ -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 diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index 44dfa240d42..83b9c34636e 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -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, diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index ce743063309..289012659a1 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -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 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts index 1a02e363de9..83847261fe8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts @@ -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; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx index 9f7e029a1d4..b1c026d3904 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx @@ -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 = ({ userRole, acces const [orgToDelete, setOrgToDelete] = useState(null); const [isDeleting, setIsDeleting] = useState(false); const [isOrgModalVisible, setIsOrgModalVisible] = useState(false); - const [form] = Form.useForm(); const [showFilters, setShowFilters] = useState(false); const [filters, setFilters] = useState({ org_id: "", org_alias: "" }); @@ -83,48 +77,6 @@ const OrganizationsPanel: React.FC = ({ 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 (
@@ -190,97 +142,7 @@ const OrganizationsPanel: React.FC = ({ userRole, acces )} - -
- - - - - form.setFieldValue("models", values)} - context="organization" - /> - - - - - - - - daily - weekly - monthly - - - - - - - - - - - Allowed Vector Stores{" "} - - - - - } - name="allowed_vector_store_ids" - className="mt-4" - help="Select vector stores this organization can access. Leave empty for access to all vector stores" - > - form.setFieldValue("allowed_vector_store_ids", values)} - value={form.getFieldValue("allowed_vector_store_ids")} - accessToken={accessToken || ""} - placeholder="Select vector stores (optional)" - /> - - - - Allowed MCP Servers{" "} - - - - - } - name="allowed_mcp_servers_and_groups" - className="mt-4" - help="Select MCP servers and access groups this organization can access." - > - 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)" - /> - - - - - - -
- -
-
-
+ ({ // Mock the child components to simplify testing vi.mock("@/components/activity_metrics", () => ({ - ActivityMetrics: () =>
Activity Metrics
, - processActivityData: () => ({ data: [], metadata: {} }), + ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => ( +
+ Activity Metrics + {`metrics-source:${modelMetrics?.__source ?? "none"}`} +
+ ), + processActivityData: (_data: unknown, key: string) => ({ __source: key }), +})); + +vi.mock("../EndpointUsage/EndpointUsage", () => ({ + default: () =>
Endpoint Usage Panel
, })); 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(); + + 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(); + + 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: [], diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 534e2be7fe8..e330983b6f9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -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 = ({ accessToken, entityType, enti const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1); + const costPanel = ( + + {/* Total Spend Card */} + + + {capitalizedEntityLabel} Spend Overview + + + Total Spend + + ${formatNumberWithCommas(spendData.metadata.total_spend, 2)} + + + + Total Requests + {spendData.metadata.total_api_requests.toLocaleString()} + + + Successful Requests + + {spendData.metadata.total_successful_requests.toLocaleString()} + + + + Failed Requests + + {spendData.metadata.total_failed_requests.toLocaleString()} + + + + Total Tokens + {spendData.metadata.total_tokens.toLocaleString()} + + + + + + {/* Daily Spend Chart */} + + + + Daily Spend + + + 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 ( +
+

{data.date}

+

Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}

+

Total Requests: {data.metrics.api_requests}

+

Successful: {data.metrics.successful_requests}

+

Failed: {data.metrics.failed_requests}

+

Total Tokens: {data.metrics.total_tokens}

+

+ Total {capitalizedEntityLabel}s: {entityCount} +

+
+

Spend by {capitalizedEntityLabel}:

+ {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 ( +

+ {getEntityLabel(entity, metrics.metadata)}: $ + {formatNumberWithCommas(metrics.metrics.spend, 2)} +

+ ); + })} + {entityCount > 5 &&

...and {entityCount - 5} more

} +
+
+ ); + }} + /> +
+
+ + + {/* Entity Breakdown Section */} + + +
+
+ Spend Per {capitalizedEntityLabel} + Showing Top 5 by Spend +
+ Get Started by Tracking cost per {capitalizedEntityLabel} + + here + +
+
+ + + { + if (!active || !payload?.[0]) return null; + const data = payload[0].payload; + return ( +
+

{data.metadata.alias}

+

Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}

+

Requests: {data.metrics.api_requests.toLocaleString()}

+

+ Successful: {data.metrics.successful_requests.toLocaleString()} +

+

Failed: {data.metrics.failed_requests.toLocaleString()}

+

Tokens: {data.metrics.total_tokens.toLocaleString()}

+
+ ); + }} + /> + + +
+ + + + {capitalizedEntityLabel} + Spend + Successful + Failed + Tokens + + + + {getEntityBreakdown() + .filter((entity) => entity.metrics.spend > 0) + .map((entity) => ( + + {entity.metadata.alias} + + + + + {entity.metrics.successful_requests.toLocaleString()} + + + {entity.metrics.failed_requests.toLocaleString()} + + {entity.metrics.total_tokens.toLocaleString()} + + ))} + +
+
+ +
+
+
+ + + {/* Top API Keys */} + + + Top Virtual Keys + + + + + {/* Top Models */} + + + {entityType === "agent" ? "Top Agents" : "Top Models"} + + + + + {/* Top Agents - only for team entity type */} + {entityType === "team" && ( + + + Top Agents Driving Spend + + + + )} + + {/* Spend by Provider */} + + +
+ Provider Usage + + + `$${formatNumberWithCommas(value, 2)}`} + colors={["cyan", "blue", "indigo", "violet", "purple"]} + showLabel + startAngle={90} + endAngle={-270} + /> + + + + + + Provider + Spend + Successful + Failed + Tokens + + + + {getProviderSpend().map((provider) => ( + + +
+ {provider.provider && } + {provider.provider} +
+
+ + + + + {provider.successful_requests.toLocaleString()} + + {provider.failed_requests.toLocaleString()} + {provider.tokens.toLocaleString()} +
+ ))} +
+
+ +
+
+
+ +
+ ); + + 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: , + }, + ...(entityType === "team" + ? [{ key: "agents", label: "Agent Activity", content: }] + : []), + { + key: "keys", + label: "Key Activity", + content: , + }, + { key: "endpoints", label: "Endpoint Activity", content: }, + ]; + return (
{isFetchingMore && ( @@ -501,320 +799,14 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti /> - Cost - {entityType === "agent" ? "Request / Token Consumption" : "Model Activity"} - {entityType === "team" ? Agent Activity : <>} - Key Activity - Endpoint Activity + {tabs.map(({ key, label }) => ( + {label} + ))} - - - {/* Total Spend Card */} - - - {capitalizedEntityLabel} Spend Overview - - - Total Spend - - ${formatNumberWithCommas(spendData.metadata.total_spend, 2)} - - - - Total Requests - - {spendData.metadata.total_api_requests.toLocaleString()} - - - - Successful Requests - - {spendData.metadata.total_successful_requests.toLocaleString()} - - - - Failed Requests - - {spendData.metadata.total_failed_requests.toLocaleString()} - - - - Total Tokens - - {spendData.metadata.total_tokens.toLocaleString()} - - - - - - - {/* Daily Spend Chart */} - - - - Daily Spend - - - 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 ( -
-

{data.date}

-

- Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)} -

-

Total Requests: {data.metrics.api_requests}

-

Successful: {data.metrics.successful_requests}

-

Failed: {data.metrics.failed_requests}

-

Total Tokens: {data.metrics.total_tokens}

-

- Total {capitalizedEntityLabel}s: {entityCount} -

-
-

Spend by {capitalizedEntityLabel}:

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

- {getEntityLabel(entity, metrics.metadata)}: $ - {formatNumberWithCommas(metrics.metrics.spend, 2)} -

- ); - })} - {entityCount > 5 && ( -

...and {entityCount - 5} more

- )} -
-
- ); - }} - /> -
-
- - - {/* Entity Breakdown Section */} - - -
-
- Spend Per {capitalizedEntityLabel} - Showing Top 5 by Spend -
- Get Started by Tracking cost per {capitalizedEntityLabel} - - here - -
-
- - - { - if (!active || !payload?.[0]) return null; - const data = payload[0].payload; - return ( -
-

{data.metadata.alias}

-

Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}

-

Requests: {data.metrics.api_requests.toLocaleString()}

-

- Successful: {data.metrics.successful_requests.toLocaleString()} -

-

Failed: {data.metrics.failed_requests.toLocaleString()}

-

Tokens: {data.metrics.total_tokens.toLocaleString()}

-
- ); - }} - /> - - -
- - - - {capitalizedEntityLabel} - Spend - Successful - Failed - Tokens - - - - {getEntityBreakdown() - .filter((entity) => entity.metrics.spend > 0) - .map((entity) => ( - - {entity.metadata.alias} - - - - - {entity.metrics.successful_requests.toLocaleString()} - - - {entity.metrics.failed_requests.toLocaleString()} - - {entity.metrics.total_tokens.toLocaleString()} - - ))} - -
-
- -
-
-
- - - {/* Top API Keys */} - - - Top Virtual Keys - - - - - {/* Top Models */} - - - {entityType === "agent" ? "Top Agents" : "Top Models"} - - - - - {/* Top Agents - only for team entity type */} - {entityType === "team" && ( - - - Top Agents Driving Spend - - - - )} - - {/* Spend by Provider */} - - -
- Provider Usage - - - `$${formatNumberWithCommas(value, 2)}`} - colors={["cyan", "blue", "indigo", "violet", "purple"]} - showLabel - startAngle={90} - endAngle={-270} - /> - - - - - - Provider - Spend - Successful - Failed - Tokens - - - - {getProviderSpend().map((provider) => ( - - -
- {provider.provider && } - {provider.provider} -
-
- - - - - {provider.successful_requests.toLocaleString()} - - - {provider.failed_requests.toLocaleString()} - - {provider.tokens.toLocaleString()} -
- ))} -
-
- -
-
-
- -
-
- - - - {entityType === "team" ? ( - - - - ) : ( - <> - )} - - - - - - + {tabs.map(({ key, content }) => ( + {content} + ))}
diff --git a/ui/litellm-dashboard/src/components/SSOModals.test.tsx b/ui/litellm-dashboard/src/components/SSOModals.test.tsx index c5da4ee7064..0fcfe60ffe8 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.test.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.test.tsx @@ -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 ( + {}} + handleAddSSOCancel={() => {}} + handleShowInstructions={mockHandleShowInstructions} + handleInstructionsOk={() => {}} + handleInstructionsCancel={() => {}} + form={form} + accessToken="test-token" + ssoConfigured={false} + /> + ); + }; + + render(); + + 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, diff --git a/ui/litellm-dashboard/src/components/SSOModals.tsx b/ui/litellm-dashboard/src/components/SSOModals.tsx index 637abbf4a81..be57e14ff60 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.tsx @@ -21,6 +21,17 @@ interface SSOModalsProps { ssoConfigured?: boolean; // Add optional prop to indicate if SSO is configured } +const detectSSOProvider = (values: Record): 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 = ({ isAddSSOModalVisible, isInstructionsModalVisible, @@ -43,22 +54,7 @@ const SSOModals: React.FC = ({ 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 = ({ 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 = ({ ...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 = ({ 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, diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx index caa6ff4f1e8..7deac9cbbbc 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx @@ -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 = { { 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 /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) => ( - - {field.name.includes("client") ? : } - - )); + 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 = ; + } else if (field.type === "textarea") { + control = ; + } else if (field.type === "password" || field.name.includes("client")) { + control = ; + } else { + control = ; + } + return ( + + {control} + + ); + }); }; const BaseSSOSettingsForm: React.FC = ({ form, onFormSubmit }) => { diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/DeleteSSOSettingsModal.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/DeleteSSOSettingsModal.tsx index 2656c861aa8..cbb55be6c1f 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/DeleteSSOSettingsModal.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/DeleteSSOSettingsModal.tsx @@ -29,6 +29,10 @@ const DeleteSSOSettingsModal: React.FC = ({ 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, diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx index d2d54033395..7415683af83 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx @@ -181,7 +181,8 @@ vi.mock("@/components/shared/errorUtils", () => ({ parseErrorMessage: vi.fn(), })); -vi.mock("../utils", () => ({ +vi.mock("../utils", async (importOriginal) => ({ + ...(await importOriginal()), processSSOSettingsPayload: vi.fn(), })); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx index 6c341c42fb7..97fbc31ce2f 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx @@ -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 = ({ 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 = ({ 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 diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx index e585bec4fd5..ed54549a40d 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.test.tsx @@ -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: "", + 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(); + }); }); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx index 849921c1076..0c83994cde6 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx @@ -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 ? ( + Provided + ) : ( + Not configured + ), + }, + { + label: "SP Entity ID", + render: (values: SSOSettingsValues) => renderEndpointValue(values.saml_sp_entity_id), + }, + { + label: "Allow IdP-initiated (unsolicited) responses", + render: (values: SSOSettingsValues) => ( + + {values.saml_allow_unsolicited === "true" ? "Enabled" : "Disabled"} + + ), + }, + { label: "Proxy Base URL", render: (values: SSOSettingsValues) => renderSimpleValue(values.proxy_base_url) }, + ], + }, }; const renderSSOSettings = () => { diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts index b5f5ccb1b8c..19a64a59ca1 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/constants.ts @@ -7,6 +7,7 @@ export const ssoProviderLogoMap: Record = { 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 = { microsoft: "Microsoft SSO", okta: "Okta / Auth0 SSO", generic: "Generic SSO", + saml: "SAML SSO", }; export const defaultRoleDisplayNames: Record = { diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.test.ts b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.test.ts index 722d52d64f9..b280f5e92ec 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.test.ts +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.test.ts @@ -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: "" } 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"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.ts b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.ts index 948ed4d2bfe..768fc7e8e0e 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.ts +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/utils.ts @@ -22,6 +22,10 @@ export const processSSOSettingsPayload = (formValues: Record): 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; }; diff --git a/ui/litellm-dashboard/src/components/agents/types.ts b/ui/litellm-dashboard/src/components/agents/types.ts index c29c566a5fe..24ff0c0e12c 100644 --- a/ui/litellm-dashboard/src/components/agents/types.ts +++ b/ui/litellm-dashboard/src/components/agents/types.ts @@ -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; -} +export type AgentObjectPermission = components["schemas"]["AgentObjectPermission"]; export interface Agent { agent_id: string; diff --git a/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.test.tsx b/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.test.tsx index 45121013652..896d8a14717 100644 --- a/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.test.tsx +++ b/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.test.tsx @@ -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) => ; - const Select = ({ children, value, onChange, placeholder }: any) => ( - - ); - Select.Option = Option; - return { - Select, - Tooltip: ({ children, title }: any) => ( -
- {children} -
- ), - Switch: ({ checked, onChange }: any) => ( - onChange(e.target.checked)} /> - ), - 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: () => ℹ, -})); +interface HarnessProps { + isCreateMode?: boolean; + onFinish?: (values: Record) => void; +} -vi.mock("@tremor/react", () => ({ - TextInput: ({ value, onValueChange, onChange, placeholder, name, className }: any) => { - const handleChange = (e: React.ChangeEvent) => { - if (onChange) { - onChange(e); - } - if (onValueChange) { - onValueChange(e.target.value); - } - }; - return ( - = ({ isCreateMode = true, onFinish = () => {} }) => { + const [form] = Form.useForm(); + const [autoRotationEnabled, setAutoRotationEnabled] = useState(false); + const [rotationInterval, setRotationInterval] = useState(""); + const [neverExpire, setNeverExpire] = useState(false); + + return ( +
+ - ); - }, -})); + + + {rotationInterval} + + ); +}; + +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(); - + it("renders the expiry and auto-rotation sections", () => { + renderWithProviders(); 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(); + it("uses the create-mode placeholder in create mode", () => { + renderWithProviders(); + 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(); + expect(screen.getByPlaceholderText(EDIT_PLACEHOLDER)).toBeInTheDocument(); + }); - it("should show correct placeholder in create mode", () => { - renderWithProviders(); - - 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(); - - const input = screen.getByTestId("duration-input"); - expect(input).toHaveAttribute("placeholder", "e.g., 30d"); - }); - - it("should show correct tooltip in create mode", () => { - renderWithProviders(); - - 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(); - - 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(); - - 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(); + const onFinish = vi.fn(); + renderWithProviders(); - 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(); + renderWithProviders(); - 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(); + + // 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(); - - expect(screen.getByText("Enable Auto-Rotation")).toBeInTheDocument(); - expect(screen.getByTestId("switch")).toBeInTheDocument(); - }); - - it("should show switch as unchecked when autoRotationEnabled is false", () => { - renderWithProviders(); - - const switchElement = screen.getByTestId("switch") as HTMLInputElement; - expect(switchElement.checked).toBe(false); - }); - - it("should show switch as checked when autoRotationEnabled is true", () => { - renderWithProviders(); - - 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(); + const onFinish = vi.fn(); + renderWithProviders(); - 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(); + describe("Auto-Rotation", () => { + it("reveals the rotation interval controls when enabled", async () => { + const user = userEvent.setup(); + renderWithProviders(); 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(); - - expect(screen.getByText("Rotation Interval")).toBeInTheDocument(); - expect(screen.getByTestId("select")).toBeInTheDocument(); - }); - - it("should show all predefined interval options", () => { - renderWithProviders(); - - 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(); - - 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( - , - ); + renderWithProviders(); - 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(); + renderWithProviders(); - 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( - , - ); + renderWithProviders(); - 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( - , - ); - - 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(); - - 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(); - - 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(); - - 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(); - - 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( - , - ); + renderWithProviders(); - 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(); }); }); }); diff --git a/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx b/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx index 7c4738f9ede..8e88fab1095 100644 --- a/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx +++ b/ui/litellm-dashboard/src/components/common_components/KeyLifecycleSettings.tsx @@ -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 = ({ const [showCustomInput, setShowCustomInput] = useState(isCustomInterval); const [customInterval, setCustomInterval] = useState(isCustomInterval ? rotationInterval : ""); - const [durationValue, setDurationValue] = useState(form?.getFieldValue?.("duration") || ""); const handleIntervalChange = (value: string) => { if (value === "custom") { @@ -53,14 +52,6 @@ const KeyLifecycleSettings: React.FC = ({ 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 (
{/* Key Expiry Section */} @@ -80,7 +71,6 @@ const KeyLifecycleSettings: React.FC = ({ 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 = ({ )} - + + +
diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx b/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx index e1c1fcb232c..4b446b0c283 100644 --- a/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx +++ b/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx @@ -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; - vector_stores: string[]; - agents?: string[]; - agent_access_groups?: string[]; - }; + object_permission?: ObjectPermission | null; access_group_ids?: string[]; budget_fallbacks?: Record; budget_limits?: Array<{ budget_duration: string; max_budget: number; reset_at?: string }>; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 051b83f4e27..576e16cbb37 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -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, // 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, // Assuming formValues is an object diff --git a/ui/litellm-dashboard/src/components/object_permission_types.ts b/ui/litellm-dashboard/src/components/object_permission_types.ts new file mode 100644 index 00000000000..bde7281faec --- /dev/null +++ b/ui/litellm-dashboard/src/components/object_permission_types.ts @@ -0,0 +1,3 @@ +import type { components } from "@/lib/http/schema"; + +export type ObjectPermission = Partial; diff --git a/ui/litellm-dashboard/src/components/object_permissions_view.tsx b/ui/litellm-dashboard/src/components/object_permissions_view.tsx index be021a5b59d..687d1a5a846 100644 --- a/ui/litellm-dashboard/src/components/object_permissions_view.tsx +++ b/ui/litellm-dashboard/src/components/object_permissions_view.tsx @@ -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; - 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; diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx index f84b95b8d4d..2f72fa37d86 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx @@ -443,6 +443,36 @@ describe("CreateKey", () => { }); }); + it("should include mcp_toolsets in keyCreateCall payload when only toolsets are selected", async () => { + renderWithProviders(); + + 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( = ({ 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 = ({ team, teams, data, addKey, autoOp /> - diff --git a/ui/litellm-dashboard/src/components/organization/org-create/OrgCreateDialog.test.tsx b/ui/litellm-dashboard/src/components/organization/org-create/OrgCreateDialog.test.tsx new file mode 100644 index 00000000000..5ec9bb1e633 --- /dev/null +++ b/ui/litellm-dashboard/src/components/organization/org-create/OrgCreateDialog.test.tsx @@ -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 }) => ( + + ), +})); +vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ + __esModule: true, + default: ({ onChange }: { onChange: (values: string[]) => void }) => ( + + ), +})); +vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({ + __esModule: true, + default: ({ + onChange, + }: { + onChange: (values: { servers: string[]; accessGroups: string[]; toolsets: string[] }) => void; + }) => ( + + ), +})); + +import { OrgCreateDialog } from "./OrgCreateDialog"; + +const Harness = ({ createOrganization }: { createOrganization: (body: unknown) => Promise }) => { + const [open, setOpen] = React.useState(true); + return ( + <> + + + + ); +}; + +const renderDialog = (overrides?: { createOrganization?: ReturnType }) => { + const createOrganization = overrides?.createOrganization ?? vi.fn().mockResolvedValue({}); + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + render( + + + , + ); + 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()); + }); +}); diff --git a/ui/litellm-dashboard/src/components/organization/org-create/OrgCreateDialog.tsx b/ui/litellm-dashboard/src/components/organization/org-create/OrgCreateDialog.tsx new file mode 100644 index 00000000000..998d9446365 --- /dev/null +++ b/ui/litellm-dashboard/src/components/organization/org-create/OrgCreateDialog.tsx @@ -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 => { + const { data } = await fetchClient.POST("/organization/new", { body }); + return data; +}; + +interface OrgCreateDialogProps { + open: boolean; + onOpenChange: (open: boolean) => void; + accessToken: string; + createOrganization?: (body: OrgCreateBody) => Promise; +} + +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 ( + + + + Create Organization + + +
+ + + {({ ref, ...field }) => } + + + + {(field) => ( + + )} + + + + {({ ref, ...field }) => } + + + + {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( + + )} + + + + {({ ref, ...field }) => } + + + + {({ ref, ...field }) => } + + + + {(field) => ( + + )} + + + + {(field) => ( + + )} + + + + {({ ref, ...field }) =>