mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
chore: merge litellm_internal_staging into litellm_fix_responses_bridge_streaming_contract
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
f68abdc861
105 changed files with 5765 additions and 1900 deletions
2
.github/workflows/_test-unit-base.yml
vendored
2
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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'
|
||||
|
|
|
|||
2
.github/workflows/mutation-test.yml
vendored
2
.github/workflows/mutation-test.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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 \
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
493
litellm/proxy/management_endpoints/sso/saml_sso.py
Normal file
493
litellm/proxy/management_endpoints/sso/saml_sso.py
Normal file
|
|
@ -0,0 +1,493 @@
|
|||
"""
|
||||
SAML 2.0 SSO for the LiteLLM proxy admin UI.
|
||||
|
||||
Supports both SP-initiated and IdP-initiated login via the HTTP-POST binding,
|
||||
using the OneLogin python3-saml toolkit for signature, audience and time
|
||||
validation. The IdP is configured from its metadata (``SAML_IDP_METADATA_URL``
|
||||
or inline ``SAML_IDP_METADATA_XML``); a successful login is mapped to a
|
||||
``CustomOpenID`` and handed to the shared post-login path used by every other
|
||||
SSO provider.
|
||||
|
||||
python3-saml pulls in the native ``xmlsec``/``libxml2`` libraries, so it is an
|
||||
optional dependency. When it is not installed the SAML routes return a clear
|
||||
error instead of breaking proxy startup.
|
||||
"""
|
||||
|
||||
# python3-saml ships no type stubs, so the type checker sees every onelogin call
|
||||
# as Unknown and the guarded optional import as possibly-unbound. Values crossing
|
||||
# that boundary are cast() to concrete types at each use site; these directives
|
||||
# silence only the unavoidable noise from the untyped dependency in this module.
|
||||
# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false
|
||||
# pyright: reportUnknownArgumentType=false, reportUnknownParameterType=false
|
||||
# pyright: reportMissingTypeStubs=false, reportPossiblyUnboundVariable=false
|
||||
# pyright: reportConstantRedefinition=false
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import secrets
|
||||
import time
|
||||
from typing import cast
|
||||
from urllib.parse import parse_qsl
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from fastapi.responses import RedirectResponse
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.management_endpoints.types import CustomOpenID, get_litellm_user_role
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
|
||||
try:
|
||||
from onelogin.saml2.auth import OneLogin_Saml2_Auth
|
||||
from onelogin.saml2.idp_metadata_parser import OneLogin_Saml2_IdPMetadataParser
|
||||
from onelogin.saml2.settings import OneLogin_Saml2_Settings
|
||||
from onelogin.saml2.xml_utils import OneLogin_Saml2_XML
|
||||
|
||||
SAML_AVAILABLE = True
|
||||
except ImportError:
|
||||
SAML_AVAILABLE = False
|
||||
|
||||
SAML_LOGIN_ROUTE = "sso/saml/login"
|
||||
SAML_CALLBACK_ROUTE = "sso/saml/callback"
|
||||
SAML_METADATA_ROUTE = "sso/saml/metadata"
|
||||
|
||||
_SAML_AUTHN_STATE_COOKIE = "litellm_saml_authn"
|
||||
_SAML_IDP_SETTINGS_CACHE_PREFIX = "saml_idp_settings"
|
||||
_SAML_AUTHN_REQUEST_CACHE_PREFIX = "saml_authn_request"
|
||||
_SAML_CONSUMED_ASSERTION_CACHE_PREFIX = "saml_consumed_assertion"
|
||||
_SAML_AUTHN_REQUEST_TTL_SECONDS = 600
|
||||
_SAML_IDP_METADATA_TTL_SECONDS = 3600
|
||||
_SAML_METADATA_FETCH_TIMEOUT_SECONDS = 10
|
||||
_SAML_MAX_POST_BYTES = 5 * 1024 * 1024
|
||||
# The replay guard tracks each assertion's NotOnOrAfter so it spans the full
|
||||
# validity window; the floor covers IdPs that issue hour-long assertions or omit
|
||||
# the timestamp, and the cap bounds cache growth.
|
||||
_SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS = 3600
|
||||
_SAML_REPLAY_GUARD_MAX_TTL_SECONDS = 86400
|
||||
|
||||
_EMAIL_ATTRIBUTE_CANDIDATES = (
|
||||
"urn:oid:0.9.2342.19200300.100.1.3",
|
||||
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/emailaddress",
|
||||
"email",
|
||||
"emailAddress",
|
||||
"mail",
|
||||
"Email",
|
||||
)
|
||||
_FIRST_NAME_ATTRIBUTE_CANDIDATES = (
|
||||
"urn:oid:2.5.4.42",
|
||||
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/givenname",
|
||||
"givenName",
|
||||
"first_name",
|
||||
"firstName",
|
||||
)
|
||||
_LAST_NAME_ATTRIBUTE_CANDIDATES = (
|
||||
"urn:oid:2.5.4.4",
|
||||
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/surname",
|
||||
"sn",
|
||||
"surname",
|
||||
"last_name",
|
||||
"lastName",
|
||||
)
|
||||
_ROLE_ATTRIBUTE_CANDIDATES = ("role", "roles", "litellm_role")
|
||||
_TEAM_IDS_ATTRIBUTE_CANDIDATES = ("teams", "team_ids", "groups")
|
||||
|
||||
|
||||
def _saml_unavailable_error() -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=status.HTTP_501_NOT_IMPLEMENTED,
|
||||
detail=(
|
||||
"SAML SSO requires the optional 'python3-saml' dependency, which is "
|
||||
"not installed. Re-install litellm with the saml extra: "
|
||||
"'pip install litellm[saml]'. The saml extra bundles the native "
|
||||
"xmlsec/libxml2 libraries, so no system packages are required."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class SAMLAuthHandler:
|
||||
"""SP- and IdP-initiated SAML 2.0 login for the admin UI."""
|
||||
|
||||
@staticmethod
|
||||
def _env(name: str, default: str | None = None) -> str | None:
|
||||
return os.getenv(name, default)
|
||||
|
||||
@staticmethod
|
||||
def is_saml_configured() -> bool:
|
||||
return bool(SAMLAuthHandler._env("SAML_IDP_METADATA_URL") or SAMLAuthHandler._env("SAML_IDP_METADATA_XML"))
|
||||
|
||||
@staticmethod
|
||||
def _bool_env(name: str, default: bool) -> bool:
|
||||
raw = SAMLAuthHandler._env(name)
|
||||
if raw is None:
|
||||
return default
|
||||
return raw.strip().lower() in ("true", "1", "yes", "on")
|
||||
|
||||
@staticmethod
|
||||
def _base_url(request: Request) -> str:
|
||||
base = get_custom_url(request_base_url=str(request.base_url))
|
||||
return base if base.endswith("/") else base + "/"
|
||||
|
||||
@staticmethod
|
||||
def _is_https(request: Request) -> bool:
|
||||
return SAMLAuthHandler._base_url(request).startswith("https")
|
||||
|
||||
@staticmethod
|
||||
def _acs_url(request: Request) -> str:
|
||||
return SAMLAuthHandler._base_url(request) + SAML_CALLBACK_ROUTE
|
||||
|
||||
@staticmethod
|
||||
def _metadata_url(request: Request) -> str:
|
||||
return SAMLAuthHandler._base_url(request) + SAML_METADATA_ROUTE
|
||||
|
||||
@staticmethod
|
||||
def _sp_entity_id(request: Request) -> str:
|
||||
return SAMLAuthHandler._env("SAML_SP_ENTITY_ID") or SAMLAuthHandler._metadata_url(request)
|
||||
|
||||
@staticmethod
|
||||
async def _load_idp_settings(cache: DualCache) -> dict[str, object]:
|
||||
metadata_url = SAMLAuthHandler._env("SAML_IDP_METADATA_URL")
|
||||
metadata_xml = SAMLAuthHandler._env("SAML_IDP_METADATA_XML")
|
||||
source = metadata_url or metadata_xml
|
||||
if source is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_501_NOT_IMPLEMENTED,
|
||||
detail="SAML SSO is not configured. Set SAML_IDP_METADATA_URL or SAML_IDP_METADATA_XML.",
|
||||
)
|
||||
|
||||
cache_key = f"{_SAML_IDP_SETTINGS_CACHE_PREFIX}:{hashlib.sha256(source.encode()).hexdigest()}"
|
||||
cached = cache.get_cache(key=cache_key)
|
||||
if isinstance(cached, dict):
|
||||
return cast(dict[str, object], cached) # cast-ok: untyped python3-saml
|
||||
|
||||
if metadata_url is not None:
|
||||
parsed = await asyncio.to_thread(
|
||||
OneLogin_Saml2_IdPMetadataParser.parse_remote,
|
||||
metadata_url,
|
||||
validate_cert=SAMLAuthHandler._bool_env("SAML_IDP_METADATA_VALIDATE_CERT", True),
|
||||
timeout=_SAML_METADATA_FETCH_TIMEOUT_SECONDS,
|
||||
)
|
||||
else:
|
||||
parsed = OneLogin_Saml2_IdPMetadataParser.parse(cast(str, metadata_xml)) # cast-ok: untyped python3-saml
|
||||
|
||||
idp_settings = cast(dict[str, object], parsed) # cast-ok: untyped python3-saml
|
||||
if not idp_settings.get("idp"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="Could not parse an IdP entityID/SSO URL/certificate from the SAML metadata.",
|
||||
)
|
||||
cache.set_cache(key=cache_key, value=idp_settings, ttl=_SAML_IDP_METADATA_TTL_SECONDS)
|
||||
return idp_settings
|
||||
|
||||
@staticmethod
|
||||
def _build_settings(request: Request, idp_settings: dict[str, object]) -> dict[str, object]:
|
||||
sp_settings: dict[str, object] = {
|
||||
"strict": SAMLAuthHandler._bool_env("SAML_STRICT", True),
|
||||
"debug": False,
|
||||
"sp": {
|
||||
"entityId": SAMLAuthHandler._sp_entity_id(request),
|
||||
"assertionConsumerService": {
|
||||
"url": SAMLAuthHandler._acs_url(request),
|
||||
"binding": "urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST",
|
||||
},
|
||||
"NameIDFormat": SAMLAuthHandler._env(
|
||||
"SAML_SP_NAME_ID_FORMAT",
|
||||
"urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress",
|
||||
),
|
||||
},
|
||||
"security": {
|
||||
"wantAssertionsSigned": SAMLAuthHandler._bool_env("SAML_WANT_ASSERTIONS_SIGNED", True),
|
||||
"wantMessagesSigned": SAMLAuthHandler._bool_env("SAML_WANT_MESSAGES_SIGNED", False),
|
||||
"authnRequestsSigned": SAMLAuthHandler._bool_env("SAML_AUTHN_REQUESTS_SIGNED", False),
|
||||
"wantNameId": True,
|
||||
"requestedAuthnContext": False,
|
||||
"rejectUnsolicitedResponsesWithInResponseTo": False,
|
||||
},
|
||||
}
|
||||
return OneLogin_Saml2_IdPMetadataParser.merge_settings(sp_settings, idp_settings)
|
||||
|
||||
@staticmethod
|
||||
def _prepare_request_data(request: Request, post_data: dict[str, str] | None = None) -> dict[str, object]:
|
||||
base = SAMLAuthHandler._base_url(request)
|
||||
scheme, _, host_part = base.partition("://")
|
||||
host = host_part.split("/", 1)[0]
|
||||
return {
|
||||
"https": "on" if scheme == "https" else "off",
|
||||
"http_host": host,
|
||||
"script_name": "/" + SAML_CALLBACK_ROUTE,
|
||||
"get_data": dict(request.query_params),
|
||||
"post_data": post_data or {},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
async def _build_auth(
|
||||
request: Request,
|
||||
cache: DualCache,
|
||||
post_data: dict[str, str] | None = None,
|
||||
) -> "OneLogin_Saml2_Auth":
|
||||
if not SAML_AVAILABLE:
|
||||
raise _saml_unavailable_error()
|
||||
idp_settings = await SAMLAuthHandler._load_idp_settings(cache)
|
||||
settings = SAMLAuthHandler._build_settings(request, idp_settings)
|
||||
request_data = SAMLAuthHandler._prepare_request_data(request, post_data)
|
||||
try:
|
||||
return OneLogin_Saml2_Auth(request_data, old_settings=settings)
|
||||
except Exception as e: # noqa: BLE001 - toolkit exposes no common exception base; fail closed
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Invalid SAML configuration: {e}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def build_login_redirect(
|
||||
request: Request, cache: DualCache, relay_state: str | None = None
|
||||
) -> RedirectResponse:
|
||||
auth = await SAMLAuthHandler._build_auth(request, cache)
|
||||
redirect_url = cast(str, auth.login(return_to=relay_state)) # cast-ok: untyped python3-saml
|
||||
response = RedirectResponse(url=redirect_url, status_code=303)
|
||||
request_id = cast(str | None, auth.get_last_request_id()) # cast-ok: untyped python3-saml
|
||||
if request_id is not None:
|
||||
cache.set_cache(
|
||||
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}",
|
||||
value="1",
|
||||
ttl=_SAML_AUTHN_REQUEST_TTL_SECONDS,
|
||||
)
|
||||
secure = SAMLAuthHandler._is_https(request)
|
||||
response.set_cookie(
|
||||
key=_SAML_AUTHN_STATE_COOKIE,
|
||||
value=request_id,
|
||||
max_age=_SAML_AUTHN_REQUEST_TTL_SECONDS,
|
||||
httponly=True,
|
||||
secure=secure,
|
||||
samesite="none" if secure else "lax",
|
||||
)
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
async def build_sp_metadata(request: Request, cache: DualCache) -> str:
|
||||
if not SAML_AVAILABLE:
|
||||
raise _saml_unavailable_error()
|
||||
idp_settings = await SAMLAuthHandler._load_idp_settings(cache)
|
||||
settings = SAMLAuthHandler._build_settings(request, idp_settings)
|
||||
saml_settings = OneLogin_Saml2_Settings(settings, sp_validation_only=True)
|
||||
metadata = cast(str, saml_settings.get_sp_metadata()) # cast-ok: untyped python3-saml
|
||||
errors = cast(list[str], saml_settings.validate_metadata(metadata)) # cast-ok: untyped python3-saml
|
||||
if errors:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Invalid SP metadata: {', '.join(errors)}",
|
||||
)
|
||||
return metadata
|
||||
|
||||
@staticmethod
|
||||
async def read_acs_post_data(request: Request) -> dict[str, str]:
|
||||
"""Read the ACS POST form under a hard size cap before any base64/XML decoding.
|
||||
|
||||
Bounds both Content-Length-declared and chunked requests so an unauthenticated
|
||||
caller cannot force unbounded buffering while decoding the SAMLResponse."""
|
||||
declared = request.headers.get("content-length")
|
||||
if declared is not None and declared.isdigit() and int(declared) > _SAML_MAX_POST_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_CONTENT_TOO_LARGE,
|
||||
detail="SAML response exceeds the maximum allowed size.",
|
||||
)
|
||||
|
||||
body = bytearray()
|
||||
async for chunk in request.stream():
|
||||
body += chunk
|
||||
if len(body) > _SAML_MAX_POST_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_CONTENT_TOO_LARGE,
|
||||
detail="SAML response exceeds the maximum allowed size.",
|
||||
)
|
||||
|
||||
return dict(parse_qsl(body.decode("utf-8", "replace")))
|
||||
|
||||
@staticmethod
|
||||
async def handle_acs(request: Request, cache: DualCache, post_data: dict[str, str]) -> CustomOpenID:
|
||||
auth = await SAMLAuthHandler._build_auth(request, cache, post_data=post_data)
|
||||
browser_request_id = request.cookies.get(_SAML_AUTHN_STATE_COOKIE)
|
||||
try:
|
||||
auth.process_response(request_id=browser_request_id)
|
||||
except Exception as e: # noqa: BLE001 - toolkit exposes no common exception base; fail closed
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=f"Could not process SAML response: {e}",
|
||||
)
|
||||
|
||||
errors = cast(list[str], auth.get_errors()) # cast-ok: untyped python3-saml
|
||||
if errors or not auth.is_authenticated():
|
||||
reason = auth.get_last_error_reason()
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=f"SAML authentication failed: {reason or ', '.join(errors)}",
|
||||
)
|
||||
|
||||
await SAMLAuthHandler._enforce_response_binding(auth, cache, browser_request_id)
|
||||
return SAMLAuthHandler._result_from_auth(auth)
|
||||
|
||||
@staticmethod
|
||||
def _replay_guard_ttl(auth: "OneLogin_Saml2_Auth") -> int:
|
||||
not_on_or_after = auth.get_last_assertion_not_on_or_after()
|
||||
if not isinstance(not_on_or_after, int):
|
||||
return _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS
|
||||
remaining = not_on_or_after - int(time.time())
|
||||
return min(
|
||||
max(remaining, _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS),
|
||||
_SAML_REPLAY_GUARD_MAX_TTL_SECONDS,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _response_in_response_to(auth: "OneLogin_Saml2_Auth") -> str | None:
|
||||
"""The request id this response answers, read from the Response element or, when the
|
||||
IdP only stamps it on the bearer SubjectConfirmationData, from there. A non-None value
|
||||
marks the response as solicited (SP-initiated) and so requiring browser binding."""
|
||||
value = cast(str | None, auth.get_last_response_in_response_to()) # cast-ok: untyped python3-saml
|
||||
if value:
|
||||
return value
|
||||
xml = cast(bytes | None, auth.get_last_response_xml()) # cast-ok: untyped python3-saml
|
||||
if not xml:
|
||||
return None
|
||||
root = OneLogin_Saml2_XML.to_etree(xml)
|
||||
for node in OneLogin_Saml2_XML.query(root, "//saml:SubjectConfirmationData[@InResponseTo]"):
|
||||
irt = cast(str | None, node.get("InResponseTo")) # cast-ok: untyped python3-saml
|
||||
if irt:
|
||||
return irt
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _enforce_response_binding(
|
||||
auth: "OneLogin_Saml2_Auth",
|
||||
cache: DualCache,
|
||||
browser_request_id: str | None,
|
||||
) -> None:
|
||||
in_response_to = SAMLAuthHandler._response_in_response_to(auth)
|
||||
|
||||
if in_response_to is not None:
|
||||
authn_key = f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{in_response_to}"
|
||||
if cache.get_cache(key=authn_key) is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="SAML response references an unknown or already-used login request.",
|
||||
)
|
||||
if browser_request_id is None or not secrets.compare_digest(browser_request_id, in_response_to):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="SAML response is not bound to this browser's login request.",
|
||||
)
|
||||
elif browser_request_id is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="SAML response is not bound to this browser's login request.",
|
||||
)
|
||||
elif not SAMLAuthHandler._bool_env("SAML_ALLOW_UNSOLICITED", False):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Unsolicited (IdP-initiated) SAML responses are disabled.",
|
||||
)
|
||||
elif cache.redis_cache is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=(
|
||||
"Unsolicited (IdP-initiated) SAML responses require a shared Redis cache "
|
||||
"so the replay guard is enforced across every worker."
|
||||
),
|
||||
)
|
||||
|
||||
assertion_id = cast(str | None, auth.get_last_assertion_id()) # cast-ok: untyped python3-saml
|
||||
if assertion_id is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="SAML assertion is missing the required ID attribute.",
|
||||
)
|
||||
consumed_key = f"{_SAML_CONSUMED_ASSERTION_CACHE_PREFIX}:{assertion_id}"
|
||||
consumed_count = await cache.async_increment_cache(
|
||||
key=consumed_key, value=1, ttl=SAMLAuthHandler._replay_guard_ttl(auth)
|
||||
)
|
||||
if consumed_count is not None and consumed_count > 1:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="SAML assertion has already been used (replay detected).",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _result_from_auth(auth: "OneLogin_Saml2_Auth") -> CustomOpenID:
|
||||
attributes = cast(dict[str, list[str]], auth.get_attributes()) # cast-ok: untyped python3-saml
|
||||
name_id = cast(str | None, auth.get_nameid()) # cast-ok: untyped python3-saml
|
||||
|
||||
email = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_EMAIL", _EMAIL_ATTRIBUTE_CANDIDATES)
|
||||
if email is None and name_id is not None and "@" in name_id:
|
||||
email = name_id
|
||||
|
||||
if email is None and SAMLAuthHandler._env("ALLOWED_EMAIL_DOMAINS") is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=(
|
||||
"SAML assertion did not contain an email address, but ALLOWED_EMAIL_DOMAINS "
|
||||
"restricts sign-in by email domain."
|
||||
),
|
||||
)
|
||||
|
||||
user_id = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_USER_ID", ()) or name_id or email
|
||||
if user_id is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="SAML assertion did not contain a usable subject (NameID) or email.",
|
||||
)
|
||||
|
||||
first_name = SAMLAuthHandler._attribute_value(
|
||||
attributes, "SAML_ATTRIBUTE_FIRST_NAME", _FIRST_NAME_ATTRIBUTE_CANDIDATES
|
||||
)
|
||||
last_name = SAMLAuthHandler._attribute_value(
|
||||
attributes, "SAML_ATTRIBUTE_LAST_NAME", _LAST_NAME_ATTRIBUTE_CANDIDATES
|
||||
)
|
||||
role_value = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_ROLE", _ROLE_ATTRIBUTE_CANDIDATES)
|
||||
team_ids = SAMLAuthHandler._attribute_values(
|
||||
attributes, "SAML_ATTRIBUTE_TEAM_IDS", _TEAM_IDS_ATTRIBUTE_CANDIDATES
|
||||
)
|
||||
|
||||
display_name = " ".join(part for part in (first_name, last_name) if part) or email
|
||||
|
||||
verbose_proxy_logger.info(f"SAML login: subject={user_id}, email={email}, attributes={list(attributes.keys())}")
|
||||
|
||||
try:
|
||||
return CustomOpenID(
|
||||
id=user_id,
|
||||
email=email,
|
||||
first_name=first_name,
|
||||
last_name=last_name,
|
||||
display_name=display_name,
|
||||
picture=None,
|
||||
provider="saml",
|
||||
team_ids=team_ids,
|
||||
user_role=get_litellm_user_role(role_value) if role_value else None,
|
||||
)
|
||||
except ValidationError as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=f"SAML assertion contained an invalid subject or email: {e}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _attribute_value(
|
||||
attributes: dict[str, list[str]],
|
||||
env_override: str,
|
||||
candidates: tuple[str, ...],
|
||||
) -> str | None:
|
||||
values = SAMLAuthHandler._attribute_values(attributes, env_override, candidates)
|
||||
return values[0] if values else None
|
||||
|
||||
@staticmethod
|
||||
def _attribute_values(
|
||||
attributes: dict[str, list[str]],
|
||||
env_override: str,
|
||||
candidates: tuple[str, ...],
|
||||
) -> list[str]:
|
||||
override = SAMLAuthHandler._env(env_override)
|
||||
keys = (override, *candidates) if override else candidates
|
||||
for key in keys:
|
||||
values = attributes.get(key)
|
||||
if values:
|
||||
return [v for v in values if v]
|
||||
return []
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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'",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
58
tests/e2e/mcp/test_mcp_access_group_e2e.py
Normal file
58
tests/e2e/mcp/test_mcp_access_group_e2e.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
"""Live e2e: MCP tool selection via access group at key creation.
|
||||
|
||||
An admin registers the Datadog remote MCP server tagged with a server-side
|
||||
access group (`mcp_access_groups`). A key minted with that access group
|
||||
(`object_permission.mcp_access_groups`) sees the server's tools; a key minted
|
||||
with a different group does not. This exercises access-group-scoped tool
|
||||
selection, the enterprise MCP surface where keys are granted tool access groups
|
||||
rather than explicit server ids.
|
||||
|
||||
A tools/list that leaks the server across the access-group boundary fails hard.
|
||||
Requires DD_API_KEY + DD_APP_KEY (the suite's real MCP upstream).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from mcp_client import McpClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
class TestMcpAccessGroupToolSelection:
|
||||
@pytest.mark.covers("mcp.list_tools.api_key.access_group_scoped")
|
||||
def test_access_group_scopes_tool_selection(
|
||||
self, client: McpClient, resources: ResourceManager
|
||||
) -> None:
|
||||
group = f"e2e-mcp-grp-{unique_marker()}"
|
||||
server_id = register_datadog_mcp(client, resources, mcp_access_groups=[group])
|
||||
|
||||
granted = client.generate_key(
|
||||
user_id=f"e2e-mcp-ag-granted-{unique_marker()}",
|
||||
mcp_servers=None,
|
||||
mcp_access_groups=[group],
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(granted))
|
||||
|
||||
other = client.generate_key(
|
||||
user_id=f"e2e-mcp-ag-other-{unique_marker()}",
|
||||
mcp_servers=None,
|
||||
mcp_access_groups=[f"e2e-mcp-grp-absent-{unique_marker()}"],
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(other))
|
||||
|
||||
granted_tools = unwrap(client.list_tools(granted))
|
||||
assert granted_tools.tool_name_containing(server_id, SEARCH_LOGS_TOOL) is not None, (
|
||||
f"key granted access group {group} did not see the tagged server's tool "
|
||||
f"(upstream dead or access-group grant not applied): "
|
||||
f"{granted_tools.tool_names_for_server(server_id)}"
|
||||
)
|
||||
|
||||
other_tools = unwrap(client.list_tools(other)).tool_names_for_server(server_id)
|
||||
assert other_tools == frozenset(), (
|
||||
f"key with a different access group saw the server's tools; access-group tool "
|
||||
f"selection leaked across the boundary: {other_tools}"
|
||||
)
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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) == []
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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({}) == ""
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
644
tests/test_litellm/proxy/management_endpoints/test_saml_sso.py
Normal file
644
tests/test_litellm/proxy/management_endpoints/test_saml_sso.py
Normal file
|
|
@ -0,0 +1,644 @@
|
|||
"""
|
||||
Regression tests for SAML 2.0 SSO (SP- and IdP-initiated) on the admin UI.
|
||||
|
||||
These exercise the real OneLogin python3-saml validation by generating signed
|
||||
SAML responses with a freshly minted IdP keypair, so a mutation that weakens
|
||||
signature, signing-requirement, expiry, replay or attribute-mapping handling
|
||||
makes a test fail.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import datetime
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
pytest.importorskip(
|
||||
"onelogin", reason="python3-saml (saml extra) is required for SAML SSO tests"
|
||||
)
|
||||
|
||||
from cryptography import x509
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from cryptography.x509.oid import NameOID
|
||||
from onelogin.saml2.utils import OneLogin_Saml2_Utils
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from typing import cast
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.management_endpoints.sso.saml_sso import (
|
||||
_SAML_AUTHN_REQUEST_CACHE_PREFIX,
|
||||
_SAML_AUTHN_STATE_COOKIE,
|
||||
_SAML_MAX_POST_BYTES,
|
||||
_SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS,
|
||||
_SAML_REPLAY_GUARD_MAX_TTL_SECONDS,
|
||||
SAMLAuthHandler,
|
||||
)
|
||||
|
||||
|
||||
def _shared_cache(store=None):
|
||||
"""A DualCache whose replay guard is backed by a shared, atomic store.
|
||||
|
||||
An InMemoryCache instance stands in for Redis; passing the same instance to
|
||||
two DualCaches simulates two workers sharing one atomic backend."""
|
||||
return DualCache(redis_cache=cast(RedisCache, store or InMemoryCache()))
|
||||
|
||||
IDP_ENTITY = "https://idp.example.com/metadata"
|
||||
SP_ENTITY = "https://proxy.example.com/sso/saml/metadata"
|
||||
ACS = "https://proxy.example.com/sso/saml/callback"
|
||||
SSO_URL = "https://idp.example.com/sso"
|
||||
PROXY_BASE_URL = "https://proxy.example.com"
|
||||
|
||||
|
||||
def _make_idp_keypair():
|
||||
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "idp.example.com")])
|
||||
cert = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(name)
|
||||
.issuer_name(name)
|
||||
.public_key(key.public_key())
|
||||
.serial_number(x509.random_serial_number())
|
||||
.not_valid_before(datetime.datetime.utcnow() - datetime.timedelta(days=1))
|
||||
.not_valid_after(datetime.datetime.utcnow() + datetime.timedelta(days=365))
|
||||
.sign(key, hashes.SHA256())
|
||||
)
|
||||
key_pem = key.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.TraditionalOpenSSL,
|
||||
serialization.NoEncryption(),
|
||||
).decode()
|
||||
cert_pem = cert.public_bytes(serialization.Encoding.PEM).decode()
|
||||
return key_pem, cert_pem
|
||||
|
||||
|
||||
def _idp_metadata_xml(cert_pem):
|
||||
cert_body = "".join(
|
||||
line for line in cert_pem.splitlines() if "CERTIFICATE" not in line
|
||||
)
|
||||
return (
|
||||
'<?xml version="1.0"?>'
|
||||
f'<EntityDescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" entityID="{IDP_ENTITY}">'
|
||||
'<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">'
|
||||
'<KeyDescriptor use="signing"><KeyInfo xmlns="http://www.w3.org/2000/09/xmldsig#">'
|
||||
f"<X509Data><X509Certificate>{cert_body}</X509Certificate></X509Data>"
|
||||
"</KeyInfo></KeyDescriptor>"
|
||||
'<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" '
|
||||
f'Location="{SSO_URL}"/>'
|
||||
"</IDPSSODescriptor></EntityDescriptor>"
|
||||
)
|
||||
|
||||
|
||||
def _saml_time(delta_seconds):
|
||||
t = datetime.datetime.utcnow() + datetime.timedelta(seconds=delta_seconds)
|
||||
return t.strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
|
||||
|
||||
def _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
*,
|
||||
in_response_to=None,
|
||||
response_level_in_response_to=True,
|
||||
email="alice@example.com",
|
||||
attributes=None,
|
||||
not_before_delta=-60,
|
||||
not_on_or_after_delta=300,
|
||||
sign=True,
|
||||
):
|
||||
if attributes is None:
|
||||
attributes = {
|
||||
"email": [email],
|
||||
"givenName": ["Alice"],
|
||||
"sn": ["Smith"],
|
||||
"role": ["internal_user"],
|
||||
}
|
||||
assertion_id = "_assertion_" + OneLogin_Saml2_Utils.generate_unique_id()
|
||||
response_id = "_response_" + OneLogin_Saml2_Utils.generate_unique_id()
|
||||
not_before = _saml_time(not_before_delta)
|
||||
not_on_or_after = _saml_time(not_on_or_after_delta)
|
||||
issue_instant = _saml_time(-1)
|
||||
irt = f'InResponseTo="{in_response_to}"' if in_response_to else ""
|
||||
response_irt = irt if response_level_in_response_to else ""
|
||||
|
||||
attr_xml = "".join(
|
||||
f'<saml:Attribute Name="{name}">'
|
||||
+ "".join(f"<saml:AttributeValue>{v}</saml:AttributeValue>" for v in values)
|
||||
+ "</saml:Attribute>"
|
||||
for name, values in attributes.items()
|
||||
)
|
||||
|
||||
assertion = (
|
||||
'<saml:Assertion xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion" '
|
||||
f'ID="{assertion_id}" Version="2.0" IssueInstant="{issue_instant}">'
|
||||
f"<saml:Issuer>{IDP_ENTITY}</saml:Issuer>"
|
||||
"<saml:Subject>"
|
||||
'<saml:NameID Format="urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress">'
|
||||
f"{email}</saml:NameID>"
|
||||
'<saml:SubjectConfirmation Method="urn:oasis:names:tc:SAML:2.0:cm:bearer">'
|
||||
f'<saml:SubjectConfirmationData {irt} NotOnOrAfter="{not_on_or_after}" Recipient="{ACS}"/>'
|
||||
"</saml:SubjectConfirmation></saml:Subject>"
|
||||
f'<saml:Conditions NotBefore="{not_before}" NotOnOrAfter="{not_on_or_after}">'
|
||||
f"<saml:AudienceRestriction><saml:Audience>{SP_ENTITY}</saml:Audience>"
|
||||
"</saml:AudienceRestriction></saml:Conditions>"
|
||||
f'<saml:AuthnStatement AuthnInstant="{issue_instant}" SessionIndex="_session">'
|
||||
"<saml:AuthnContext><saml:AuthnContextClassRef>"
|
||||
"urn:oasis:names:tc:SAML:2.0:ac:classes:Password"
|
||||
"</saml:AuthnContextClassRef></saml:AuthnContext></saml:AuthnStatement>"
|
||||
f"<saml:AttributeStatement>{attr_xml}</saml:AttributeStatement>"
|
||||
"</saml:Assertion>"
|
||||
)
|
||||
|
||||
if sign:
|
||||
signed = OneLogin_Saml2_Utils.add_sign(assertion, key_pem, cert_pem)
|
||||
assertion = (signed.decode() if isinstance(signed, bytes) else signed).replace(
|
||||
'<?xml version="1.0"?>', ""
|
||||
)
|
||||
|
||||
return (
|
||||
'<?xml version="1.0"?>'
|
||||
'<samlp:Response xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol" '
|
||||
'xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion" '
|
||||
f'ID="{response_id}" Version="2.0" IssueInstant="{issue_instant}" '
|
||||
f'Destination="{ACS}" {response_irt}>'
|
||||
f"<saml:Issuer>{IDP_ENTITY}</saml:Issuer>"
|
||||
'<samlp:Status><samlp:StatusCode Value="urn:oasis:names:tc:SAML:2.0:status:Success"/>'
|
||||
"</samlp:Status>"
|
||||
f"{assertion}</samlp:Response>"
|
||||
)
|
||||
|
||||
|
||||
def _b64(xml):
|
||||
return base64.b64encode(xml.encode()).decode()
|
||||
|
||||
|
||||
def _fake_request(cookies=None):
|
||||
return type(
|
||||
"Req",
|
||||
(),
|
||||
{
|
||||
"base_url": URL(PROXY_BASE_URL + "/"),
|
||||
"query_params": {},
|
||||
"cookies": cookies or {},
|
||||
},
|
||||
)()
|
||||
|
||||
|
||||
async def _acs(b64, cache, cookies=None):
|
||||
return await SAMLAuthHandler.handle_acs(
|
||||
_fake_request(cookies), cache, {"SAMLResponse": b64}
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def saml_env(monkeypatch):
|
||||
key_pem, cert_pem = _make_idp_keypair()
|
||||
monkeypatch.setenv("SAML_IDP_METADATA_XML", _idp_metadata_xml(cert_pem))
|
||||
monkeypatch.setenv("SAML_SP_ENTITY_ID", SP_ENTITY)
|
||||
monkeypatch.setenv("PROXY_BASE_URL", PROXY_BASE_URL)
|
||||
for var in (
|
||||
"SAML_IDP_METADATA_URL",
|
||||
"SAML_ATTRIBUTE_EMAIL",
|
||||
"SAML_ATTRIBUTE_TEAM_IDS",
|
||||
"SAML_ALLOW_UNSOLICITED",
|
||||
"ALLOWED_EMAIL_DOMAINS",
|
||||
):
|
||||
monkeypatch.delenv(var, raising=False)
|
||||
return key_pem, cert_pem
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def saml_env_idp_initiated(saml_env, monkeypatch):
|
||||
monkeypatch.setenv("SAML_ALLOW_UNSOLICITED", "true")
|
||||
return saml_env
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_valid_idp_initiated_login_maps_assertion_to_user(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
resp = _build_signed_response(key_pem, cert_pem)
|
||||
|
||||
result = await _acs(_b64(resp), _shared_cache())
|
||||
|
||||
assert result.email == "alice@example.com"
|
||||
assert result.id == "alice@example.com"
|
||||
assert result.first_name == "Alice"
|
||||
assert result.last_name == "Smith"
|
||||
assert result.user_role == LitellmUserRoles.INTERNAL_USER
|
||||
assert result.provider == "saml"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tampered_assertion_is_rejected(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
resp = _build_signed_response(key_pem, cert_pem)
|
||||
tampered = resp.replace("alice@example.com", "attacker@example.com")
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(tampered), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsigned_assertion_is_rejected(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
resp = _build_signed_response(key_pem, cert_pem, sign=False)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_signature_from_untrusted_key_is_rejected(saml_env):
|
||||
_, cert_pem = saml_env
|
||||
attacker_key, attacker_cert = _make_idp_keypair()
|
||||
resp = _build_signed_response(attacker_key, attacker_cert)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_expired_assertion_is_rejected(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
resp = _build_signed_response(
|
||||
key_pem, cert_pem, not_before_delta=-7200, not_on_or_after_delta=-3600
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sp_initiated_unknown_in_response_to_is_rejected(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
resp = _build_signed_response(key_pem, cert_pem, in_response_to="_never_issued")
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sp_initiated_known_request_succeeds_once_then_replay_rejected(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
cache = DualCache()
|
||||
request_id = "_authn_req_known"
|
||||
cache.set_cache(
|
||||
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
|
||||
)
|
||||
resp = _build_signed_response(key_pem, cert_pem, in_response_to=request_id)
|
||||
cookies = {_SAML_AUTHN_STATE_COOKIE: request_id}
|
||||
|
||||
result = await _acs(_b64(resp), cache, cookies=cookies)
|
||||
assert result.email == "alice@example.com"
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), cache, cookies=cookies)
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sp_initiated_response_not_bound_to_browser_is_rejected(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
cache = DualCache()
|
||||
request_id = "_authn_req_known"
|
||||
cache.set_cache(
|
||||
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
|
||||
)
|
||||
resp = _build_signed_response(key_pem, cert_pem, in_response_to=request_id)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), cache)
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(
|
||||
_b64(resp), cache, cookies={_SAML_AUTHN_STATE_COOKIE: "_attacker_request"}
|
||||
)
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subjectconfirmation_only_in_response_to_without_cookie_is_rejected(
|
||||
saml_env_idp_initiated,
|
||||
):
|
||||
"""An IdP that stamps InResponseTo only on the SubjectConfirmationData (not the
|
||||
Response element) is still solicited and must be browser-bound: with unsolicited
|
||||
explicitly allowed, a missing cookie must still 401 rather than slip through."""
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
cache = DualCache()
|
||||
request_id = "_authn_req_known"
|
||||
cache.set_cache(
|
||||
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
|
||||
)
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
in_response_to=request_id,
|
||||
response_level_in_response_to=False,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), cache)
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subjectconfirmation_only_in_response_to_with_cookie_succeeds(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
cache = DualCache()
|
||||
request_id = "_authn_req_known"
|
||||
cache.set_cache(
|
||||
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
|
||||
)
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
in_response_to=request_id,
|
||||
response_level_in_response_to=False,
|
||||
)
|
||||
|
||||
result = await _acs(
|
||||
_b64(resp), cache, cookies={_SAML_AUTHN_STATE_COOKIE: request_id}
|
||||
)
|
||||
assert result.email == "alice@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsolicited_response_rejected_by_default(saml_env):
|
||||
key_pem, cert_pem = saml_env
|
||||
resp = _build_signed_response(key_pem, cert_pem)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_idp_initiated_assertion_replay_is_rejected(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
cache = _shared_cache()
|
||||
resp = _build_signed_response(key_pem, cert_pem, email="bob@example.com")
|
||||
|
||||
first = await _acs(_b64(resp), cache)
|
||||
assert first.email == "bob@example.com"
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), cache)
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_assertion_without_id_is_rejected(saml_env_idp_initiated):
|
||||
"""An assertion with no ID attribute has no stable replay key. On the unsolicited
|
||||
path there is no browser binding, so the consumed-assertion guard is the only replay
|
||||
defense; a missing ID must be rejected rather than silently skipping the guard."""
|
||||
|
||||
class _AuthNoAssertionId:
|
||||
def get_last_response_in_response_to(self):
|
||||
return None
|
||||
|
||||
def get_last_response_xml(self):
|
||||
return None
|
||||
|
||||
def get_last_assertion_id(self):
|
||||
return None
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await SAMLAuthHandler._enforce_response_binding(
|
||||
_AuthNoAssertionId(), _shared_cache(), None
|
||||
)
|
||||
assert exc.value.status_code == 401
|
||||
assert "ID" in exc.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsolicited_response_rejected_when_disabled(saml_env, monkeypatch):
|
||||
key_pem, cert_pem = saml_env
|
||||
monkeypatch.setenv("SAML_ALLOW_UNSOLICITED", "false")
|
||||
resp = _build_signed_response(key_pem, cert_pem)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_email_in_assertion_is_rejected_cleanly(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
email="not-an-email",
|
||||
attributes={"email": ["not-an-email"], "givenName": ["X"]},
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), _shared_cache())
|
||||
assert exc.value.status_code == 401
|
||||
assert "invalid subject or email" in exc.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_email_less_assertion_rejected_when_domain_restriction_configured(
|
||||
saml_env_idp_initiated, monkeypatch
|
||||
):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
monkeypatch.setenv("ALLOWED_EMAIL_DOMAINS", "example.com")
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
email="opaque-persistent-id-123",
|
||||
attributes={"givenName": ["Alice"]},
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), _shared_cache())
|
||||
assert exc.value.status_code == 401
|
||||
assert "ALLOWED_EMAIL_DOMAINS" in exc.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_email_less_assertion_allowed_without_domain_restriction(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
email="opaque-persistent-id-123",
|
||||
attributes={"givenName": ["Alice"]},
|
||||
)
|
||||
|
||||
result = await _acs(_b64(resp), _shared_cache())
|
||||
assert result.email is None
|
||||
assert result.id == "opaque-persistent-id-123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_email_attribute_override(saml_env_idp_initiated, monkeypatch):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
monkeypatch.setenv("SAML_ATTRIBUTE_EMAIL", "corpMail")
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
email="ignored@example.com",
|
||||
attributes={
|
||||
"corpMail": ["real@corp.example.com"],
|
||||
"givenName": ["Real"],
|
||||
},
|
||||
)
|
||||
|
||||
result = await _acs(_b64(resp), _shared_cache())
|
||||
assert result.email == "real@corp.example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_ids_extracted_from_groups_attribute(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
resp = _build_signed_response(
|
||||
key_pem,
|
||||
cert_pem,
|
||||
attributes={
|
||||
"email": ["carol@example.com"],
|
||||
"groups": ["team-a", "team-b"],
|
||||
},
|
||||
)
|
||||
|
||||
result = await _acs(_b64(resp), _shared_cache())
|
||||
assert result.team_ids == ["team-a", "team-b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_login_redirect_targets_idp_and_caches_request_id(saml_env):
|
||||
cache = DualCache()
|
||||
redirect = await SAMLAuthHandler.build_login_redirect(_fake_request(), cache)
|
||||
|
||||
location = redirect.headers["location"]
|
||||
assert location.startswith(SSO_URL)
|
||||
assert "SAMLRequest=" in location
|
||||
cached = [
|
||||
k
|
||||
for k in cache.in_memory_cache.cache_dict
|
||||
if k.startswith(_SAML_AUTHN_REQUEST_CACHE_PREFIX)
|
||||
]
|
||||
assert len(cached) == 1
|
||||
|
||||
request_id = cached[0].split(":", 1)[1]
|
||||
set_cookie = redirect.headers["set-cookie"]
|
||||
assert f"{_SAML_AUTHN_STATE_COOKIE}={request_id}" in set_cookie
|
||||
assert "httponly" in set_cookie.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sp_metadata_contains_acs_and_entity_id(saml_env):
|
||||
metadata = await SAMLAuthHandler.build_sp_metadata(_fake_request(), DualCache())
|
||||
assert ACS in metadata
|
||||
assert SP_ENTITY in metadata
|
||||
assert "AssertionConsumerService" in metadata
|
||||
|
||||
|
||||
def test_replay_guard_ttl_tracks_assertion_validity():
|
||||
class _Auth:
|
||||
def __init__(self, not_on_or_after):
|
||||
self._not_on_or_after = not_on_or_after
|
||||
|
||||
def get_last_assertion_not_on_or_after(self):
|
||||
return self._not_on_or_after
|
||||
|
||||
now = int(time.time())
|
||||
|
||||
long_lived = SAMLAuthHandler._replay_guard_ttl(_Auth(now + 7200))
|
||||
assert long_lived >= 7200
|
||||
|
||||
short_lived = SAMLAuthHandler._replay_guard_ttl(_Auth(now + 60))
|
||||
assert short_lived == _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS
|
||||
|
||||
missing = SAMLAuthHandler._replay_guard_ttl(_Auth(None))
|
||||
assert missing == _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS
|
||||
|
||||
capped = SAMLAuthHandler._replay_guard_ttl(_Auth(now + 10 * 86400))
|
||||
assert capped == _SAML_REPLAY_GUARD_MAX_TTL_SECONDS
|
||||
|
||||
|
||||
def test_is_saml_configured_reflects_env(monkeypatch):
|
||||
monkeypatch.delenv("SAML_IDP_METADATA_URL", raising=False)
|
||||
monkeypatch.delenv("SAML_IDP_METADATA_XML", raising=False)
|
||||
assert SAMLAuthHandler.is_saml_configured() is False
|
||||
|
||||
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml")
|
||||
assert SAMLAuthHandler.is_saml_configured() is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_idp_initiated_rejected_without_shared_cache(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
resp = _build_signed_response(key_pem, cert_pem)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), DualCache())
|
||||
assert exc.value.status_code == 401
|
||||
assert "shared Redis cache" in exc.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_idp_initiated_replay_rejected_across_workers(saml_env_idp_initiated):
|
||||
key_pem, cert_pem = saml_env_idp_initiated
|
||||
shared_store = InMemoryCache()
|
||||
worker_one = _shared_cache(shared_store)
|
||||
worker_two = _shared_cache(shared_store)
|
||||
resp = _build_signed_response(key_pem, cert_pem, email="bob@example.com")
|
||||
|
||||
first = await _acs(_b64(resp), worker_one)
|
||||
assert first.email == "bob@example.com"
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _acs(_b64(resp), worker_two)
|
||||
assert exc.value.status_code == 401
|
||||
|
||||
|
||||
class _FakeChunkedRequest:
|
||||
def __init__(self, chunks, content_length=None):
|
||||
self._chunks = chunks
|
||||
self.headers = {} if content_length is None else {"content-length": content_length}
|
||||
|
||||
async def stream(self):
|
||||
for chunk in self._chunks:
|
||||
yield chunk
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_acs_post_data_parses_form():
|
||||
body = b"SAMLResponse=abc123&RelayState=%2Fui%2F"
|
||||
request = _FakeChunkedRequest([body], content_length=str(len(body)))
|
||||
|
||||
post_data = await SAMLAuthHandler.read_acs_post_data(cast(Request, request))
|
||||
|
||||
assert post_data == {"SAMLResponse": "abc123", "RelayState": "/ui/"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_acs_post_data_rejects_oversized_content_length():
|
||||
request = _FakeChunkedRequest([b""], content_length=str(_SAML_MAX_POST_BYTES + 1))
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await SAMLAuthHandler.read_acs_post_data(cast(Request, request))
|
||||
assert exc.value.status_code == 413
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_acs_post_data_rejects_oversized_stream_without_content_length():
|
||||
chunk = b"a" * (1024 * 1024)
|
||||
chunk_count = _SAML_MAX_POST_BYTES // len(chunk) + 2
|
||||
request = _FakeChunkedRequest([chunk] * chunk_count)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await SAMLAuthHandler.read_acs_post_data(cast(Request, request))
|
||||
assert exc.value.status_code == 413
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -1,19 +1,14 @@
|
|||
import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters";
|
||||
import { InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { Form, Input, Modal, Select as Select2, Tooltip } from "antd";
|
||||
import { useQueryClient } from "@tanstack/react-query";
|
||||
import React, { useState } from "react";
|
||||
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
|
||||
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
|
||||
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
|
||||
import NotificationsManager from "@/components/molecules/notifications_manager";
|
||||
import { organizationCreateCall, organizationDeleteCall } from "@/components/networking";
|
||||
import { organizationDeleteCall } from "@/components/networking";
|
||||
import { OrgCreateDialog } from "@/components/organization/org-create/OrgCreateDialog";
|
||||
import OrganizationInfoView from "@/components/organization/organization_view";
|
||||
import NumericalInput from "@/components/shared/numerical_input";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector";
|
||||
|
||||
import OrganizationsTable from "./OrganizationsTable";
|
||||
|
||||
|
|
@ -30,7 +25,6 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
|
|||
const [orgToDelete, setOrgToDelete] = useState<string | null>(null);
|
||||
const [isDeleting, setIsDeleting] = useState(false);
|
||||
const [isOrgModalVisible, setIsOrgModalVisible] = useState(false);
|
||||
const [form] = Form.useForm();
|
||||
const [showFilters, setShowFilters] = useState(false);
|
||||
const [filters, setFilters] = useState<FilterState>({ org_id: "", org_alias: "" });
|
||||
|
||||
|
|
@ -83,48 +77,6 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
|
|||
setOrgToDelete(null);
|
||||
};
|
||||
|
||||
const handleCreate = async (values: any) => {
|
||||
try {
|
||||
if (!accessToken) return;
|
||||
|
||||
// Transform allowed_vector_store_ids and allowed_mcp_servers_and_groups into object_permission
|
||||
if (
|
||||
(values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) ||
|
||||
(values.allowed_mcp_servers_and_groups &&
|
||||
(values.allowed_mcp_servers_and_groups.servers?.length > 0 ||
|
||||
values.allowed_mcp_servers_and_groups.accessGroups?.length > 0))
|
||||
) {
|
||||
values.object_permission = {};
|
||||
if (values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) {
|
||||
values.object_permission.vector_stores = values.allowed_vector_store_ids;
|
||||
delete values.allowed_vector_store_ids;
|
||||
}
|
||||
if (values.allowed_mcp_servers_and_groups) {
|
||||
if (values.allowed_mcp_servers_and_groups.servers?.length > 0) {
|
||||
values.object_permission.mcp_servers = values.allowed_mcp_servers_and_groups.servers;
|
||||
}
|
||||
if (values.allowed_mcp_servers_and_groups.accessGroups?.length > 0) {
|
||||
values.object_permission.mcp_access_groups = values.allowed_mcp_servers_and_groups.accessGroups;
|
||||
}
|
||||
delete values.allowed_mcp_servers_and_groups;
|
||||
}
|
||||
}
|
||||
|
||||
await organizationCreateCall(accessToken, values);
|
||||
NotificationsManager.success("Organization created successfully");
|
||||
setIsOrgModalVisible(false);
|
||||
form.resetFields();
|
||||
await refetchOrganizations();
|
||||
} catch (error) {
|
||||
console.error("Error creating organization:", error);
|
||||
}
|
||||
};
|
||||
|
||||
const handleCancel = () => {
|
||||
setIsOrgModalVisible(false);
|
||||
form.resetFields();
|
||||
};
|
||||
|
||||
if (!premiumUser) {
|
||||
return (
|
||||
<div className="mx-4 mt-4">
|
||||
|
|
@ -190,97 +142,7 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
|
|||
</>
|
||||
)}
|
||||
|
||||
<Modal title="Create Organization" visible={isOrgModalVisible} width={800} footer={null} onCancel={handleCancel}>
|
||||
<Form form={form} onFinish={handleCreate} labelCol={{ span: 8 }} wrapperCol={{ span: 16 }} labelAlign="left">
|
||||
<Form.Item
|
||||
label="Organization Name"
|
||||
name="organization_alias"
|
||||
rules={[
|
||||
{
|
||||
required: true,
|
||||
message: "Please input an organization name",
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input placeholder="" />
|
||||
</Form.Item>
|
||||
<Form.Item label="Models" name="models">
|
||||
<ModelSelect
|
||||
options={{ showAllProxyModelsOverride: true, includeSpecialOptions: true }}
|
||||
value={form.getFieldValue("models")}
|
||||
onChange={(values) => form.setFieldValue("models", values)}
|
||||
context="organization"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item label="Max Budget (USD)" name="max_budget">
|
||||
<NumericalInput step={0.01} precision={2} width={200} />
|
||||
</Form.Item>
|
||||
<Form.Item label="Reset Budget" name="budget_duration">
|
||||
<Select2 defaultValue={null} placeholder="n/a">
|
||||
<Select2.Option value="24h">daily</Select2.Option>
|
||||
<Select2.Option value="7d">weekly</Select2.Option>
|
||||
<Select2.Option value="30d">monthly</Select2.Option>
|
||||
</Select2>
|
||||
</Form.Item>
|
||||
<Form.Item label="Tokens per minute Limit (TPM)" name="tpm_limit">
|
||||
<NumericalInput step={1} width={400} />
|
||||
</Form.Item>
|
||||
<Form.Item label="Requests per minute Limit (RPM)" name="rpm_limit">
|
||||
<NumericalInput step={1} width={400} />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Allowed Vector Stores{" "}
|
||||
<Tooltip title="Select which vector stores this organization can access by default. Leave empty for access to all vector stores">
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="allowed_vector_store_ids"
|
||||
className="mt-4"
|
||||
help="Select vector stores this organization can access. Leave empty for access to all vector stores"
|
||||
>
|
||||
<VectorStoreSelector
|
||||
onChange={(values) => form.setFieldValue("allowed_vector_store_ids", values)}
|
||||
value={form.getFieldValue("allowed_vector_store_ids")}
|
||||
accessToken={accessToken || ""}
|
||||
placeholder="Select vector stores (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Allowed MCP Servers{" "}
|
||||
<Tooltip title="Select which MCP servers and access groups this organization can access by default.">
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="allowed_mcp_servers_and_groups"
|
||||
className="mt-4"
|
||||
help="Select MCP servers and access groups this organization can access."
|
||||
>
|
||||
<MCPServerSelector
|
||||
onChange={(values) => form.setFieldValue("allowed_mcp_servers_and_groups", values)}
|
||||
value={form.getFieldValue("allowed_mcp_servers_and_groups")}
|
||||
accessToken={accessToken || ""}
|
||||
placeholder="Select MCP servers and access groups (optional)"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item label="Metadata" name="metadata">
|
||||
<Input.TextArea rows={4} />
|
||||
</Form.Item>
|
||||
|
||||
<div style={{ textAlign: "right", marginTop: "10px" }}>
|
||||
<Button type="submit">Create Organization</Button>
|
||||
</div>
|
||||
</Form>
|
||||
</Modal>
|
||||
<OrgCreateDialog open={isOrgModalVisible} onOpenChange={setIsOrgModalVisible} accessToken={accessToken || ""} />
|
||||
|
||||
<DeleteResourceModal
|
||||
isOpen={isDeleteModalOpen}
|
||||
|
|
|
|||
|
|
@ -25,8 +25,17 @@ vi.mock("@/components/networking", () => ({
|
|||
|
||||
// Mock the child components to simplify testing
|
||||
vi.mock("@/components/activity_metrics", () => ({
|
||||
ActivityMetrics: () => <div>Activity Metrics</div>,
|
||||
processActivityData: () => ({ data: [], metadata: {} }),
|
||||
ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => (
|
||||
<div>
|
||||
<span>Activity Metrics</span>
|
||||
<span>{`metrics-source:${modelMetrics?.__source ?? "none"}`}</span>
|
||||
</div>
|
||||
),
|
||||
processActivityData: (_data: unknown, key: string) => ({ __source: key }),
|
||||
}));
|
||||
|
||||
vi.mock("../EndpointUsage/EndpointUsage", () => ({
|
||||
default: () => <div>Endpoint Usage Panel</div>,
|
||||
}));
|
||||
|
||||
vi.mock("@/components/UsagePage/components/EntityUsage/TopKeyView", () => ({
|
||||
|
|
@ -481,6 +490,54 @@ describe("EntityUsage", () => {
|
|||
expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const selectedPanels = (container: HTMLElement) =>
|
||||
Array.from(container.querySelectorAll("div.tremor-TabPanel-root")).filter(
|
||||
(panel) => panel.getAttribute("aria-selected") === "true",
|
||||
);
|
||||
|
||||
it.each([
|
||||
["Cost", "Tag Spend Overview"],
|
||||
["Model Activity", "metrics-source:models"],
|
||||
["Key Activity", "metrics-source:api_keys"],
|
||||
["Endpoint Activity", "Endpoint Usage Panel"],
|
||||
])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => {
|
||||
const { container } = render(<EntityUsage {...defaultProps} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockTagDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText(tabLabel));
|
||||
});
|
||||
|
||||
const selected = selectedPanels(container);
|
||||
expect(selected).toHaveLength(1);
|
||||
expect(selected[0].textContent).toContain(marker);
|
||||
});
|
||||
|
||||
it.each([
|
||||
["Cost", "Team Spend Overview"],
|
||||
["Model Activity", "metrics-source:models"],
|
||||
["Agent Activity", "metrics-source:entities"],
|
||||
["Key Activity", "metrics-source:api_keys"],
|
||||
["Endpoint Activity", "Endpoint Usage Panel"],
|
||||
])("shows only the %s panel for the team entity type", async (tabLabel, marker) => {
|
||||
const { container } = render(<EntityUsage {...defaultProps} entityType="team" />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockTeamDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByText(tabLabel));
|
||||
});
|
||||
|
||||
const selected = selectedPanels(container);
|
||||
expect(selected).toHaveLength(1);
|
||||
expect(selected[0].textContent).toContain(marker);
|
||||
});
|
||||
|
||||
it("should handle empty data gracefully", async () => {
|
||||
const emptyData = {
|
||||
results: [],
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ import {
|
|||
} from "@tremor/react";
|
||||
import { ExportOutlined, LoadingOutlined } from "@ant-design/icons";
|
||||
import { Alert, Button } from "antd";
|
||||
import React, { useMemo, useState } from "react";
|
||||
import React, { type ReactNode, useMemo, useState } from "react";
|
||||
import TeamMultiSelect from "@/components/common_components/team_multi_select";
|
||||
import { ActivityMetrics, processActivityData } from "@/components/activity_metrics";
|
||||
import { UsageExportHeader } from "@/components/EntityUsageExport";
|
||||
|
|
@ -406,6 +406,304 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
|
||||
const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1);
|
||||
|
||||
const costPanel = (
|
||||
<Grid numItems={2} className="gap-2 w-full">
|
||||
{/* Total Spend Card */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>{capitalizedEntityLabel} Spend Overview</Title>
|
||||
<Grid numItems={5} className="gap-4 mt-4">
|
||||
<Card>
|
||||
<Title>Total Spend</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_api_requests.toLocaleString()}</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Successful Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-green-600">
|
||||
{spendData.metadata.total_successful_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Failed Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-red-600">
|
||||
{spendData.metadata.total_failed_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Tokens</Title>
|
||||
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_tokens.toLocaleString()}</Text>
|
||||
</Card>
|
||||
</Grid>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Daily Spend Chart */}
|
||||
<Col numColSpan={2}>
|
||||
<ShadcnCard>
|
||||
<CardHeader>
|
||||
<CardTitle className="text-base font-semibold">Daily Spend</CardTitle>
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={[...spendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime())}
|
||||
index="date"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
yAxisWidth={100}
|
||||
showLegend={false}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
const entityCount = Object.keys(data.breakdown.entities || {}).length;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.date}</p>
|
||||
<p className="text-cyan-500">Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}</p>
|
||||
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
|
||||
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
|
||||
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
|
||||
<p className="text-gray-600">Total Tokens: {data.metrics.total_tokens}</p>
|
||||
<p className="text-gray-600">
|
||||
Total {capitalizedEntityLabel}s: {entityCount}
|
||||
</p>
|
||||
<div className="mt-2 border-t pt-2">
|
||||
<p className="font-semibold">Spend by {capitalizedEntityLabel}:</p>
|
||||
{Object.entries(data.breakdown.entities || {})
|
||||
.sort(([, a], [, b]) => {
|
||||
const spendA = (a as EntityMetrics).metrics.spend;
|
||||
const spendB = (b as EntityMetrics).metrics.spend;
|
||||
return spendB - spendA;
|
||||
})
|
||||
.slice(0, 5)
|
||||
.map(([entity, entityData]) => {
|
||||
const metrics = entityData as EntityMetrics;
|
||||
return (
|
||||
<p key={entity} className="text-sm text-gray-600">
|
||||
{getEntityLabel(entity, metrics.metadata)}: $
|
||||
{formatNumberWithCommas(metrics.metrics.spend, 2)}
|
||||
</p>
|
||||
);
|
||||
})}
|
||||
{entityCount > 5 && <p className="text-sm text-gray-500 italic">...and {entityCount - 5} more</p>}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</CardContent>
|
||||
</ShadcnCard>
|
||||
</Col>
|
||||
|
||||
{/* Entity Breakdown Section */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<div className="flex flex-col space-y-2">
|
||||
<Title>Spend Per {capitalizedEntityLabel}</Title>
|
||||
<Subtitle className="text-xs">Showing Top 5 by Spend</Subtitle>
|
||||
<div className="flex items-center text-sm text-gray-500">
|
||||
<span>Get Started by Tracking cost per {capitalizedEntityLabel} </span>
|
||||
<a
|
||||
href="https://docs.litellm.ai/docs/proxy/enterprise#spend-tracking"
|
||||
className="text-blue-500 hover:text-blue-700 ml-1"
|
||||
>
|
||||
here
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
<Grid numItems={2} className="gap-6">
|
||||
<Col numColSpan={1}>
|
||||
<BarChart
|
||||
className="mt-4 h-52"
|
||||
data={getProcessedEntityBreakdownForChart()}
|
||||
index="metadata.alias_display"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
layout="vertical"
|
||||
showLegend={false}
|
||||
yAxisWidth={150}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.metadata.alias}</p>
|
||||
<p className="text-cyan-500">Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}</p>
|
||||
<p className="text-gray-600">Requests: {data.metrics.api_requests.toLocaleString()}</p>
|
||||
<p className="text-green-600">
|
||||
Successful: {data.metrics.successful_requests.toLocaleString()}
|
||||
</p>
|
||||
<p className="text-red-600">Failed: {data.metrics.failed_requests.toLocaleString()}</p>
|
||||
<p className="text-gray-600">Tokens: {data.metrics.total_tokens.toLocaleString()}</p>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<div className="h-52 overflow-y-auto">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>{capitalizedEntityLabel}</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getEntityBreakdown()
|
||||
.filter((entity) => entity.metrics.spend > 0)
|
||||
.map((entity) => (
|
||||
<TableRow key={entity.metadata.id}>
|
||||
<TableCell>{entity.metadata.alias}</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={entity.metrics.spend} decimals={4} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{entity.metrics.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">
|
||||
{entity.metrics.failed_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell>{entity.metrics.total_tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</div>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top API Keys */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>Top Virtual Keys</Title>
|
||||
<TopKeyView
|
||||
topKeys={getTopAPIKeys()}
|
||||
teams={null}
|
||||
showTags={entityType === "tag"}
|
||||
topKeysLimit={topKeysLimit}
|
||||
setTopKeysLimit={setTopKeysLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Models */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
|
||||
<TopModelView
|
||||
topModels={getTopModels()}
|
||||
topModelsLimit={topModelsLimit}
|
||||
setTopModelsLimit={setTopModelsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Agents - only for team entity type */}
|
||||
{entityType === "team" && (
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>Top Agents Driving Spend</Title>
|
||||
<TopModelView
|
||||
topModels={getTopAgents()}
|
||||
topModelsLimit={topAgentsLimit}
|
||||
setTopModelsLimit={setTopAgentsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
)}
|
||||
|
||||
{/* Spend by Provider */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<Title>Provider Usage</Title>
|
||||
<Grid numItems={2}>
|
||||
<Col numColSpan={1}>
|
||||
<DonutChart
|
||||
className="mt-4 h-40"
|
||||
data={getProviderSpend()}
|
||||
index="provider"
|
||||
category="spend"
|
||||
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
|
||||
colors={["cyan", "blue", "indigo", "violet", "purple"]}
|
||||
showLabel
|
||||
startAngle={90}
|
||||
endAngle={-270}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>Provider</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getProviderSpend().map((provider) => (
|
||||
<TableRow key={provider.provider}>
|
||||
<TableCell>
|
||||
<div className="flex items-center space-x-2">
|
||||
{provider.provider && <Logo provider={provider.provider} className="w-4 h-4" />}
|
||||
<span>{provider.provider}</span>
|
||||
</div>
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={provider.spend} decimals={2} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{provider.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">{provider.failed_requests.toLocaleString()}</TableCell>
|
||||
<TableCell>{provider.tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
</Grid>
|
||||
);
|
||||
|
||||
const tabs: readonly { key: string; label: string; content: ReactNode }[] = [
|
||||
{ key: "cost", label: "Cost", content: costPanel },
|
||||
{
|
||||
key: "models",
|
||||
label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity",
|
||||
content: <ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
|
||||
},
|
||||
...(entityType === "team"
|
||||
? [{ key: "agents", label: "Agent Activity", content: <ActivityMetrics modelMetrics={agentMetrics} /> }]
|
||||
: []),
|
||||
{
|
||||
key: "keys",
|
||||
label: "Key Activity",
|
||||
content: <ActivityMetrics modelMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
|
||||
},
|
||||
{ key: "endpoints", label: "Endpoint Activity", content: <EndpointUsage userSpendData={spendData} /> },
|
||||
];
|
||||
|
||||
return (
|
||||
<div style={{ width: "100%" }} className="relative">
|
||||
{isFetchingMore && (
|
||||
|
|
@ -501,320 +799,14 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
/>
|
||||
<TabGroup>
|
||||
<TabList variant="solid" className="mt-1">
|
||||
<Tab>Cost</Tab>
|
||||
<Tab>{entityType === "agent" ? "Request / Token Consumption" : "Model Activity"}</Tab>
|
||||
{entityType === "team" ? <Tab>Agent Activity</Tab> : <></>}
|
||||
<Tab>Key Activity</Tab>
|
||||
<Tab>Endpoint Activity</Tab>
|
||||
{tabs.map(({ key, label }) => (
|
||||
<Tab key={key}>{label}</Tab>
|
||||
))}
|
||||
</TabList>
|
||||
<TabPanels>
|
||||
<TabPanel>
|
||||
<Grid numItems={2} className="gap-2 w-full">
|
||||
{/* Total Spend Card */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>{capitalizedEntityLabel} Spend Overview</Title>
|
||||
<Grid numItems={5} className="gap-4 mt-4">
|
||||
<Card>
|
||||
<Title>Total Spend</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
{spendData.metadata.total_api_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Successful Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-green-600">
|
||||
{spendData.metadata.total_successful_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Failed Requests</Title>
|
||||
<Text className="text-2xl font-bold mt-2 text-red-600">
|
||||
{spendData.metadata.total_failed_requests.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
<Card>
|
||||
<Title>Total Tokens</Title>
|
||||
<Text className="text-2xl font-bold mt-2">
|
||||
{spendData.metadata.total_tokens.toLocaleString()}
|
||||
</Text>
|
||||
</Card>
|
||||
</Grid>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Daily Spend Chart */}
|
||||
<Col numColSpan={2}>
|
||||
<ShadcnCard>
|
||||
<CardHeader>
|
||||
<CardTitle className="text-base font-semibold">Daily Spend</CardTitle>
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
<BarChart
|
||||
data={[...spendData.results].sort(
|
||||
(a, b) => new Date(a.date).getTime() - new Date(b.date).getTime(),
|
||||
)}
|
||||
index="date"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
yAxisWidth={100}
|
||||
showLegend={false}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
const entityCount = Object.keys(data.breakdown.entities || {}).length;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.date}</p>
|
||||
<p className="text-cyan-500">
|
||||
Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}
|
||||
</p>
|
||||
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
|
||||
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
|
||||
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
|
||||
<p className="text-gray-600">Total Tokens: {data.metrics.total_tokens}</p>
|
||||
<p className="text-gray-600">
|
||||
Total {capitalizedEntityLabel}s: {entityCount}
|
||||
</p>
|
||||
<div className="mt-2 border-t pt-2">
|
||||
<p className="font-semibold">Spend by {capitalizedEntityLabel}:</p>
|
||||
{Object.entries(data.breakdown.entities || {})
|
||||
.sort(([, a], [, b]) => {
|
||||
const spendA = (a as EntityMetrics).metrics.spend;
|
||||
const spendB = (b as EntityMetrics).metrics.spend;
|
||||
return spendB - spendA;
|
||||
})
|
||||
.slice(0, 5)
|
||||
.map(([entity, entityData]) => {
|
||||
const metrics = entityData as EntityMetrics;
|
||||
return (
|
||||
<p key={entity} className="text-sm text-gray-600">
|
||||
{getEntityLabel(entity, metrics.metadata)}: $
|
||||
{formatNumberWithCommas(metrics.metrics.spend, 2)}
|
||||
</p>
|
||||
);
|
||||
})}
|
||||
{entityCount > 5 && (
|
||||
<p className="text-sm text-gray-500 italic">...and {entityCount - 5} more</p>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</CardContent>
|
||||
</ShadcnCard>
|
||||
</Col>
|
||||
|
||||
{/* Entity Breakdown Section */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<div className="flex flex-col space-y-2">
|
||||
<Title>Spend Per {capitalizedEntityLabel}</Title>
|
||||
<Subtitle className="text-xs">Showing Top 5 by Spend</Subtitle>
|
||||
<div className="flex items-center text-sm text-gray-500">
|
||||
<span>Get Started by Tracking cost per {capitalizedEntityLabel} </span>
|
||||
<a
|
||||
href="https://docs.litellm.ai/docs/proxy/enterprise#spend-tracking"
|
||||
className="text-blue-500 hover:text-blue-700 ml-1"
|
||||
>
|
||||
here
|
||||
</a>
|
||||
</div>
|
||||
</div>
|
||||
<Grid numItems={2} className="gap-6">
|
||||
<Col numColSpan={1}>
|
||||
<BarChart
|
||||
className="mt-4 h-52"
|
||||
data={getProcessedEntityBreakdownForChart()}
|
||||
index="metadata.alias_display"
|
||||
categories={["metrics.spend"]}
|
||||
colors={["cyan"]}
|
||||
valueFormatter={valueFormatterSpend}
|
||||
layout="vertical"
|
||||
showLegend={false}
|
||||
yAxisWidth={150}
|
||||
customTooltip={({ payload, active }) => {
|
||||
if (!active || !payload?.[0]) return null;
|
||||
const data = payload[0].payload;
|
||||
return (
|
||||
<div className="bg-white p-4 shadow-lg rounded-lg border">
|
||||
<p className="font-bold">{data.metadata.alias}</p>
|
||||
<p className="text-cyan-500">Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}</p>
|
||||
<p className="text-gray-600">Requests: {data.metrics.api_requests.toLocaleString()}</p>
|
||||
<p className="text-green-600">
|
||||
Successful: {data.metrics.successful_requests.toLocaleString()}
|
||||
</p>
|
||||
<p className="text-red-600">Failed: {data.metrics.failed_requests.toLocaleString()}</p>
|
||||
<p className="text-gray-600">Tokens: {data.metrics.total_tokens.toLocaleString()}</p>
|
||||
</div>
|
||||
);
|
||||
}}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<div className="h-52 overflow-y-auto">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>{capitalizedEntityLabel}</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getEntityBreakdown()
|
||||
.filter((entity) => entity.metrics.spend > 0)
|
||||
.map((entity) => (
|
||||
<TableRow key={entity.metadata.id}>
|
||||
<TableCell>{entity.metadata.alias}</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={entity.metrics.spend} decimals={4} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{entity.metrics.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">
|
||||
{entity.metrics.failed_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell>{entity.metrics.total_tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</div>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top API Keys */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>Top Virtual Keys</Title>
|
||||
<TopKeyView
|
||||
topKeys={getTopAPIKeys()}
|
||||
teams={null}
|
||||
showTags={entityType === "tag"}
|
||||
topKeysLimit={topKeysLimit}
|
||||
setTopKeysLimit={setTopKeysLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Models */}
|
||||
<Col numColSpan={1}>
|
||||
<Card>
|
||||
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
|
||||
<TopModelView
|
||||
topModels={getTopModels()}
|
||||
topModelsLimit={topModelsLimit}
|
||||
setTopModelsLimit={setTopModelsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
|
||||
{/* Top Agents - only for team entity type */}
|
||||
{entityType === "team" && (
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<Title>Top Agents Driving Spend</Title>
|
||||
<TopModelView
|
||||
topModels={getTopAgents()}
|
||||
topModelsLimit={topAgentsLimit}
|
||||
setTopModelsLimit={setTopAgentsLimit}
|
||||
/>
|
||||
</Card>
|
||||
</Col>
|
||||
)}
|
||||
|
||||
{/* Spend by Provider */}
|
||||
<Col numColSpan={2}>
|
||||
<Card>
|
||||
<div className="flex flex-col space-y-4">
|
||||
<Title>Provider Usage</Title>
|
||||
<Grid numItems={2}>
|
||||
<Col numColSpan={1}>
|
||||
<DonutChart
|
||||
className="mt-4 h-40"
|
||||
data={getProviderSpend()}
|
||||
index="provider"
|
||||
category="spend"
|
||||
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
|
||||
colors={["cyan", "blue", "indigo", "violet", "purple"]}
|
||||
showLabel
|
||||
startAngle={90}
|
||||
endAngle={-270}
|
||||
/>
|
||||
</Col>
|
||||
<Col numColSpan={1}>
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell>Provider</TableHeaderCell>
|
||||
<TableHeaderCell>Spend</TableHeaderCell>
|
||||
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
|
||||
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
|
||||
<TableHeaderCell>Tokens</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{getProviderSpend().map((provider) => (
|
||||
<TableRow key={provider.provider}>
|
||||
<TableCell>
|
||||
<div className="flex items-center space-x-2">
|
||||
{provider.provider && <Logo provider={provider.provider} className="w-4 h-4" />}
|
||||
<span>{provider.provider}</span>
|
||||
</div>
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<MoneyCell value={provider.spend} decimals={2} />
|
||||
</TableCell>
|
||||
<TableCell className="text-green-600">
|
||||
{provider.successful_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell className="text-red-600">
|
||||
{provider.failed_requests.toLocaleString()}
|
||||
</TableCell>
|
||||
<TableCell>{provider.tokens.toLocaleString()}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</Col>
|
||||
</Grid>
|
||||
</div>
|
||||
</Card>
|
||||
</Col>
|
||||
</Grid>
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
<ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />
|
||||
</TabPanel>
|
||||
{entityType === "team" ? (
|
||||
<TabPanel>
|
||||
<ActivityMetrics modelMetrics={agentMetrics} />
|
||||
</TabPanel>
|
||||
) : (
|
||||
<></>
|
||||
)}
|
||||
<TabPanel>
|
||||
<ActivityMetrics modelMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
<EndpointUsage userSpendData={spendData} />
|
||||
</TabPanel>
|
||||
{tabs.map(({ key, content }) => (
|
||||
<TabPanel key={key}>{content}</TabPanel>
|
||||
))}
|
||||
</TabPanels>
|
||||
</TabGroup>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
|
||||
import { Form } from "antd";
|
||||
import { Form, type FormInstance } from "antd";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import SSOModals from "./SSOModals";
|
||||
|
||||
|
|
@ -412,6 +412,76 @@ describe("SSOModals", () => {
|
|||
expect(mockHandleShowInstructions).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should submit SAML settings with the unsolicited toggle mapped to a 'true'/'false' string", async () => {
|
||||
const mockHandleShowInstructions = vi.fn();
|
||||
vi.mocked(updateSSOSettings).mockResolvedValue({});
|
||||
vi.mocked(getSSOSettings).mockResolvedValue({ values: {} });
|
||||
|
||||
let formInstance: FormInstance | null = null;
|
||||
|
||||
const TestWrapper = () => {
|
||||
const [form] = Form.useForm();
|
||||
formInstance = form;
|
||||
|
||||
return (
|
||||
<SSOModals
|
||||
isAddSSOModalVisible={true}
|
||||
isInstructionsModalVisible={false}
|
||||
handleAddSSOOk={() => {}}
|
||||
handleAddSSOCancel={() => {}}
|
||||
handleShowInstructions={mockHandleShowInstructions}
|
||||
handleInstructionsOk={() => {}}
|
||||
handleInstructionsCancel={() => {}}
|
||||
form={form}
|
||||
accessToken="test-token"
|
||||
ssoConfigured={false}
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
render(<TestWrapper />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(getSSOSettings).toHaveBeenCalledWith("test-token");
|
||||
});
|
||||
|
||||
formInstance?.setFieldsValue({ sso_provider: "saml" });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText("IdP Metadata URL")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Proxy Admin Email"), {
|
||||
target: { value: "admin@example.com" },
|
||||
});
|
||||
fireEvent.change(screen.getByLabelText("Proxy Base URL"), {
|
||||
target: { value: "https://proxy.example.com" },
|
||||
});
|
||||
fireEvent.change(screen.getByLabelText("IdP Metadata URL"), {
|
||||
target: { value: "https://idp.example.com/metadata" },
|
||||
});
|
||||
fireEvent.change(screen.getByLabelText("SP Entity ID"), {
|
||||
target: { value: "https://proxy.example.com/sso/saml/metadata" },
|
||||
});
|
||||
fireEvent.click(screen.getByLabelText("Allow IdP-initiated (unsolicited) responses"));
|
||||
|
||||
fireEvent.click(screen.getByText("Save"));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(updateSSOSettings).toHaveBeenCalledWith(
|
||||
"test-token",
|
||||
expect.objectContaining({
|
||||
sso_provider: "saml",
|
||||
saml_idp_metadata_url: "https://idp.example.com/metadata",
|
||||
saml_sp_entity_id: "https://proxy.example.com/sso/saml/metadata",
|
||||
saml_allow_unsolicited: "true",
|
||||
}),
|
||||
);
|
||||
});
|
||||
|
||||
expect(mockHandleShowInstructions).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should show Clear button and clear SSO settings when configured", async () => {
|
||||
const mockHandleAddSSOOk = vi.fn();
|
||||
(updateSSOSettings as any).mockResolvedValue({});
|
||||
|
|
@ -462,6 +532,10 @@ describe("SSOModals", () => {
|
|||
generic_authorization_endpoint: null,
|
||||
generic_token_endpoint: null,
|
||||
generic_userinfo_endpoint: null,
|
||||
saml_idp_metadata_url: null,
|
||||
saml_idp_metadata_xml: null,
|
||||
saml_sp_entity_id: null,
|
||||
saml_allow_unsolicited: null,
|
||||
generic_scope: null,
|
||||
proxy_base_url: null,
|
||||
user_email: null,
|
||||
|
|
|
|||
|
|
@ -21,6 +21,17 @@ interface SSOModalsProps {
|
|||
ssoConfigured?: boolean; // Add optional prop to indicate if SSO is configured
|
||||
}
|
||||
|
||||
const detectSSOProvider = (values: Record<string, unknown>): string | null => {
|
||||
if (values.google_client_id) return "google";
|
||||
if (values.microsoft_client_id) return "microsoft";
|
||||
if (values.generic_client_id) {
|
||||
const authEndpoint =
|
||||
typeof values.generic_authorization_endpoint === "string" ? values.generic_authorization_endpoint : "";
|
||||
return authEndpoint.includes("okta") || authEndpoint.includes("auth0") ? "okta" : "generic";
|
||||
}
|
||||
if (values.saml_idp_metadata_url || values.saml_idp_metadata_xml) return "saml";
|
||||
return null;
|
||||
};
|
||||
const SSOModals: React.FC<SSOModalsProps> = ({
|
||||
isAddSSOModalVisible,
|
||||
isInstructionsModalVisible,
|
||||
|
|
@ -43,22 +54,7 @@ const SSOModals: React.FC<SSOModalsProps> = ({
|
|||
const ssoData = await getSSOSettings(accessToken);
|
||||
if (ssoData && ssoData.values) {
|
||||
// Determine which SSO provider is configured
|
||||
let selectedProvider = null;
|
||||
if (ssoData.values.google_client_id) {
|
||||
selectedProvider = "google";
|
||||
} else if (ssoData.values.microsoft_client_id) {
|
||||
selectedProvider = "microsoft";
|
||||
} else if (ssoData.values.generic_client_id) {
|
||||
// Check if it looks like Okta based on endpoints
|
||||
if (
|
||||
ssoData.values.generic_authorization_endpoint?.includes("okta") ||
|
||||
ssoData.values.generic_authorization_endpoint?.includes("auth0")
|
||||
) {
|
||||
selectedProvider = "okta";
|
||||
} else {
|
||||
selectedProvider = "generic";
|
||||
}
|
||||
}
|
||||
const selectedProvider = detectSSOProvider(ssoData.values);
|
||||
|
||||
// Extract role mappings if they exist
|
||||
let roleMappingFields = {};
|
||||
|
|
@ -89,6 +85,7 @@ const SSOModals: React.FC<SSOModalsProps> = ({
|
|||
user_email: ssoData.values.user_email,
|
||||
...ssoData.values,
|
||||
...roleMappingFields,
|
||||
saml_allow_unsolicited: ssoData.values.saml_allow_unsolicited === "true",
|
||||
};
|
||||
|
||||
// Clear form first, then set values with a small delay to ensure proper initialization
|
||||
|
|
@ -129,6 +126,10 @@ const SSOModals: React.FC<SSOModalsProps> = ({
|
|||
...rest,
|
||||
};
|
||||
|
||||
if (typeof payload.saml_allow_unsolicited === "boolean") {
|
||||
payload.saml_allow_unsolicited = payload.saml_allow_unsolicited ? "true" : "false";
|
||||
}
|
||||
|
||||
// Add role mappings if use_role_mappings is checked
|
||||
if (use_role_mappings) {
|
||||
// Helper function to split comma-separated string into array
|
||||
|
|
@ -191,6 +192,10 @@ const SSOModals: React.FC<SSOModalsProps> = ({
|
|||
generic_authorization_endpoint: null,
|
||||
generic_token_endpoint: null,
|
||||
generic_userinfo_endpoint: null,
|
||||
saml_idp_metadata_url: null,
|
||||
saml_idp_metadata_xml: null,
|
||||
saml_sp_entity_id: null,
|
||||
saml_allow_unsolicited: null,
|
||||
generic_scope: null,
|
||||
proxy_base_url: null,
|
||||
user_email: null,
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ export interface SSOProviderConfig {
|
|||
name: string;
|
||||
placeholder?: string;
|
||||
required?: boolean;
|
||||
type?: "password" | "textarea" | "checkbox";
|
||||
}>;
|
||||
}
|
||||
|
||||
|
|
@ -90,6 +91,41 @@ export const ssoProviderConfigs: Record<string, SSOProviderConfig> = {
|
|||
{ label: "Scopes", name: "generic_scope", placeholder: "openid email profile", required: false },
|
||||
],
|
||||
},
|
||||
saml: {
|
||||
envVarMap: {
|
||||
saml_idp_metadata_url: "SAML_IDP_METADATA_URL",
|
||||
saml_idp_metadata_xml: "SAML_IDP_METADATA_XML",
|
||||
saml_sp_entity_id: "SAML_SP_ENTITY_ID",
|
||||
saml_allow_unsolicited: "SAML_ALLOW_UNSOLICITED",
|
||||
},
|
||||
fields: [
|
||||
{
|
||||
label: "IdP Metadata URL",
|
||||
name: "saml_idp_metadata_url",
|
||||
required: false,
|
||||
placeholder: "https://idp.example.com/metadata (use this or the metadata XML below)",
|
||||
},
|
||||
{
|
||||
label: "IdP Metadata XML",
|
||||
name: "saml_idp_metadata_xml",
|
||||
required: false,
|
||||
type: "textarea",
|
||||
placeholder: "Paste the IdP metadata XML here if you do not have a metadata URL",
|
||||
},
|
||||
{
|
||||
label: "SP Entity ID",
|
||||
name: "saml_sp_entity_id",
|
||||
required: false,
|
||||
placeholder: "Defaults to <proxy base url>/sso/saml/metadata",
|
||||
},
|
||||
{
|
||||
label: "Allow IdP-initiated (unsolicited) responses",
|
||||
name: "saml_allow_unsolicited",
|
||||
required: false,
|
||||
type: "checkbox",
|
||||
},
|
||||
],
|
||||
},
|
||||
};
|
||||
|
||||
// Helper function to render provider fields
|
||||
|
|
@ -97,16 +133,31 @@ export const renderProviderFields = (provider: string) => {
|
|||
const config = ssoProviderConfigs[provider];
|
||||
if (!config) return null;
|
||||
|
||||
return config.fields.map((field) => (
|
||||
<Form.Item
|
||||
key={field.name}
|
||||
label={field.label}
|
||||
name={field.name}
|
||||
rules={[{ required: field.required !== false, message: `Please enter the ${field.label.toLowerCase()}` }]}
|
||||
>
|
||||
{field.name.includes("client") ? <Input.Password /> : <TextInput placeholder={field.placeholder} />}
|
||||
</Form.Item>
|
||||
));
|
||||
return config.fields.map((field) => {
|
||||
const isRequired = field.required !== false;
|
||||
const rules = isRequired ? [{ required: true, message: `Please enter the ${field.label.toLowerCase()}` }] : [];
|
||||
let control: React.ReactNode;
|
||||
if (field.type === "checkbox") {
|
||||
control = <Checkbox />;
|
||||
} else if (field.type === "textarea") {
|
||||
control = <Input.TextArea rows={4} placeholder={field.placeholder} />;
|
||||
} else if (field.type === "password" || field.name.includes("client")) {
|
||||
control = <Input.Password />;
|
||||
} else {
|
||||
control = <TextInput placeholder={field.placeholder} />;
|
||||
}
|
||||
return (
|
||||
<Form.Item
|
||||
key={field.name}
|
||||
label={field.label}
|
||||
name={field.name}
|
||||
rules={rules}
|
||||
valuePropName={field.type === "checkbox" ? "checked" : undefined}
|
||||
>
|
||||
{control}
|
||||
</Form.Item>
|
||||
);
|
||||
});
|
||||
};
|
||||
|
||||
const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormSubmit }) => {
|
||||
|
|
|
|||
|
|
@ -29,6 +29,10 @@ const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({ isVisib
|
|||
generic_authorization_endpoint: null,
|
||||
generic_token_endpoint: null,
|
||||
generic_userinfo_endpoint: null,
|
||||
saml_idp_metadata_url: null,
|
||||
saml_idp_metadata_xml: null,
|
||||
saml_sp_entity_id: null,
|
||||
saml_allow_unsolicited: null,
|
||||
proxy_base_url: null,
|
||||
user_email: null,
|
||||
sso_provider: null,
|
||||
|
|
|
|||
|
|
@ -181,7 +181,8 @@ vi.mock("@/components/shared/errorUtils", () => ({
|
|||
parseErrorMessage: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("../utils", () => ({
|
||||
vi.mock("../utils", async (importOriginal) => ({
|
||||
...(await importOriginal<typeof import("../utils")>()),
|
||||
processSSOSettingsPayload: vi.fn(),
|
||||
}));
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import React, { useEffect } from "react";
|
|||
import BaseSSOSettingsForm from "./BaseSSOSettingsForm";
|
||||
import NotificationsManager from "@/components/molecules/notifications_manager";
|
||||
import { parseErrorMessage } from "@/components/shared/errorUtils";
|
||||
import { processSSOSettingsPayload } from "../utils";
|
||||
import { detectSSOProvider, processSSOSettingsPayload } from "../utils";
|
||||
import { useSSOSettings } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
import { useEditSSOSettings } from "@/app/(dashboard)/hooks/sso/useEditSSOSettings";
|
||||
|
||||
|
|
@ -26,22 +26,7 @@ const EditSSOSettingsModal: React.FC<EditSSOSettingsModalProps> = ({ isVisible,
|
|||
const ssoData = ssoSettings.data;
|
||||
|
||||
// Determine which SSO provider is configured
|
||||
let selectedProvider = null;
|
||||
if (ssoData.values.google_client_id) {
|
||||
selectedProvider = "google";
|
||||
} else if (ssoData.values.microsoft_client_id) {
|
||||
selectedProvider = "microsoft";
|
||||
} else if (ssoData.values.generic_client_id) {
|
||||
// Check if it looks like Okta based on endpoints
|
||||
if (
|
||||
ssoData.values.generic_authorization_endpoint?.includes("okta") ||
|
||||
ssoData.values.generic_authorization_endpoint?.includes("auth0")
|
||||
) {
|
||||
selectedProvider = "okta";
|
||||
} else {
|
||||
selectedProvider = "generic";
|
||||
}
|
||||
}
|
||||
const selectedProvider = detectSSOProvider(ssoData.values);
|
||||
|
||||
// Extract role mappings if they exist
|
||||
let roleMappingFields = {};
|
||||
|
|
@ -81,6 +66,9 @@ const EditSSOSettingsModal: React.FC<EditSSOSettingsModalProps> = ({ isVisible,
|
|||
...ssoData.values,
|
||||
...roleMappingFields,
|
||||
...teamMappingFields,
|
||||
...(ssoData.values.saml_allow_unsolicited != null
|
||||
? { saml_allow_unsolicited: ssoData.values.saml_allow_unsolicited === "true" }
|
||||
: {}),
|
||||
};
|
||||
|
||||
// Clear form first, then set values with a small delay to ensure proper initialization
|
||||
|
|
|
|||
|
|
@ -48,6 +48,28 @@ const googleConfiguredValues = {
|
|||
team_mappings: null,
|
||||
};
|
||||
|
||||
const samlConfiguredValues = {
|
||||
google_client_id: null,
|
||||
google_client_secret: null,
|
||||
microsoft_client_id: null,
|
||||
microsoft_client_secret: null,
|
||||
microsoft_tenant: null,
|
||||
generic_client_id: null,
|
||||
generic_client_secret: null,
|
||||
generic_authorization_endpoint: null,
|
||||
generic_token_endpoint: null,
|
||||
generic_userinfo_endpoint: null,
|
||||
proxy_base_url: "https://proxy.example.com",
|
||||
user_email: null,
|
||||
ui_access_mode: null,
|
||||
role_mappings: null,
|
||||
team_mappings: null,
|
||||
saml_idp_metadata_url: null,
|
||||
saml_idp_metadata_xml: "<EntityDescriptor/>",
|
||||
saml_sp_entity_id: "https://proxy.example.com/sso/saml/metadata",
|
||||
saml_allow_unsolicited: "true",
|
||||
};
|
||||
|
||||
describe("SSOSettings", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
|
|
@ -77,4 +99,20 @@ describe("SSOSettings", () => {
|
|||
const logo = screen.getByAltText("Google SSO logo");
|
||||
expect(logo).toHaveAttribute("src", expect.stringContaining("google.svg"));
|
||||
});
|
||||
|
||||
it("renders a SAML configuration as configured instead of the empty placeholder", () => {
|
||||
mockUseSSOSettings.mockReturnValue({
|
||||
data: { values: samlConfiguredValues },
|
||||
isLoading: false,
|
||||
refetch: vi.fn(),
|
||||
});
|
||||
|
||||
renderSSOSettings();
|
||||
|
||||
expect(screen.queryByText("No SSO Configuration Found")).not.toBeInTheDocument();
|
||||
expect(screen.getByRole("button", { name: /Edit SSO Settings/i })).toBeInTheDocument();
|
||||
expect(screen.getByText("SAML SSO")).toBeInTheDocument();
|
||||
expect(screen.getByText("https://proxy.example.com/sso/saml/metadata")).toBeInTheDocument();
|
||||
expect(screen.getByText("Enabled")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -22,10 +22,13 @@ export default function SSOSettings() {
|
|||
const [isDeleteModalVisible, setIsDeleteModalVisible] = useState(false);
|
||||
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
|
||||
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
|
||||
const isSSOConfigured =
|
||||
Boolean(ssoSettings?.values.google_client_id) ||
|
||||
Boolean(ssoSettings?.values.microsoft_client_id) ||
|
||||
Boolean(ssoSettings?.values.generic_client_id);
|
||||
const isSSOConfigured = [
|
||||
ssoSettings?.values.google_client_id,
|
||||
ssoSettings?.values.microsoft_client_id,
|
||||
ssoSettings?.values.generic_client_id,
|
||||
ssoSettings?.values.saml_idp_metadata_url,
|
||||
ssoSettings?.values.saml_idp_metadata_xml,
|
||||
].some(Boolean);
|
||||
|
||||
const selectedProvider = ssoSettings?.values ? detectSSOProvider(ssoSettings.values) : null;
|
||||
const isRoleMappingsEnabled = Boolean(ssoSettings?.values.role_mappings);
|
||||
|
|
@ -154,6 +157,37 @@ export default function SSOSettings() {
|
|||
: null,
|
||||
],
|
||||
},
|
||||
saml: {
|
||||
providerText: ssoProviderDisplayNames.saml,
|
||||
fields: [
|
||||
{
|
||||
label: "IdP Metadata URL",
|
||||
render: (values: SSOSettingsValues) => renderEndpointValue(values.saml_idp_metadata_url),
|
||||
},
|
||||
{
|
||||
label: "IdP Metadata XML",
|
||||
render: (values: SSOSettingsValues) =>
|
||||
values.saml_idp_metadata_xml ? (
|
||||
<Tag>Provided</Tag>
|
||||
) : (
|
||||
<span className="text-gray-400 italic">Not configured</span>
|
||||
),
|
||||
},
|
||||
{
|
||||
label: "SP Entity ID",
|
||||
render: (values: SSOSettingsValues) => renderEndpointValue(values.saml_sp_entity_id),
|
||||
},
|
||||
{
|
||||
label: "Allow IdP-initiated (unsolicited) responses",
|
||||
render: (values: SSOSettingsValues) => (
|
||||
<Tag color={values.saml_allow_unsolicited === "true" ? "green" : "default"}>
|
||||
{values.saml_allow_unsolicited === "true" ? "Enabled" : "Disabled"}
|
||||
</Tag>
|
||||
),
|
||||
},
|
||||
{ label: "Proxy Base URL", render: (values: SSOSettingsValues) => renderSimpleValue(values.proxy_base_url) },
|
||||
],
|
||||
},
|
||||
};
|
||||
|
||||
const renderSSOSettings = () => {
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ export const ssoProviderLogoMap: Record<string, string> = {
|
|||
microsoft: microsoftAzureLogo.src,
|
||||
okta: "https://www.okta.com/sites/default/files/Okta_Logo_BrightBlue_Medium.png",
|
||||
generic: "",
|
||||
saml: "",
|
||||
};
|
||||
|
||||
// SSO Provider display names (consistent between select dropdown and table)
|
||||
|
|
@ -15,6 +16,7 @@ export const ssoProviderDisplayNames: Record<string, string> = {
|
|||
microsoft: "Microsoft SSO",
|
||||
okta: "Okta / Auth0 SSO",
|
||||
generic: "Generic SSO",
|
||||
saml: "SAML SSO",
|
||||
};
|
||||
|
||||
export const defaultRoleDisplayNames: Record<string, string> = {
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import { processSSOSettingsPayload } from "./utils";
|
||||
import { detectSSOProvider, processSSOSettingsPayload } from "./utils";
|
||||
import { describe, it, expect } from "vitest";
|
||||
import type { SSOSettingsValues } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
|
||||
|
||||
describe("processSSOSettingsPayload", () => {
|
||||
describe("without role mappings", () => {
|
||||
|
|
@ -428,3 +429,26 @@ describe("processSSOSettingsPayload", () => {
|
|||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("detectSSOProvider with SAML", () => {
|
||||
it("returns saml when a SAML IdP metadata URL is configured", () => {
|
||||
expect(detectSSOProvider({ saml_idp_metadata_url: "https://idp.example.com/metadata" } as SSOSettingsValues)).toBe(
|
||||
"saml",
|
||||
);
|
||||
});
|
||||
|
||||
it("returns saml when only inline SAML metadata XML is configured", () => {
|
||||
expect(detectSSOProvider({ saml_idp_metadata_xml: "<EntityDescriptor/>" } as SSOSettingsValues)).toBe("saml");
|
||||
});
|
||||
});
|
||||
|
||||
describe("processSSOSettingsPayload with SAML", () => {
|
||||
it("maps the boolean allow-unsolicited toggle to a 'true'/'false' string", () => {
|
||||
expect(
|
||||
processSSOSettingsPayload({ sso_provider: "saml", saml_allow_unsolicited: true }).saml_allow_unsolicited,
|
||||
).toBe("true");
|
||||
expect(
|
||||
processSSOSettingsPayload({ sso_provider: "saml", saml_allow_unsolicited: false }).saml_allow_unsolicited,
|
||||
).toBe("false");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -22,6 +22,10 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
|
|||
...rest,
|
||||
};
|
||||
|
||||
if (typeof payload.saml_allow_unsolicited === "boolean") {
|
||||
payload.saml_allow_unsolicited = payload.saml_allow_unsolicited ? "true" : "false";
|
||||
}
|
||||
|
||||
// Add role mappings only if use_role_mappings is checked AND provider supports role mappings
|
||||
const provider = rest.sso_provider;
|
||||
const supportsRoleMappings = provider === "okta" || provider === "generic";
|
||||
|
|
@ -81,5 +85,6 @@ export const detectSSOProvider = (values: SSOSettingsValues): string | null => {
|
|||
}
|
||||
return "generic";
|
||||
}
|
||||
if (values.saml_idp_metadata_url || values.saml_idp_metadata_xml) return "saml";
|
||||
return null;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,14 +1,12 @@
|
|||
import type { components } from "@/lib/http/schema";
|
||||
|
||||
export interface AgentAttachedKey {
|
||||
token: string;
|
||||
key_alias?: string | null;
|
||||
key_name?: string | null;
|
||||
}
|
||||
|
||||
export interface AgentObjectPermission {
|
||||
mcp_servers?: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
}
|
||||
export type AgentObjectPermission = components["schemas"]["AgentObjectPermission"];
|
||||
|
||||
export interface Agent {
|
||||
agent_id: string;
|
||||
|
|
|
|||
|
|
@ -1,357 +1,214 @@
|
|||
import React, { useState } from "react";
|
||||
// eslint-disable-next-line no-restricted-imports -- exercising KeyLifecycleSettings requires hosting it in a real antd Form (the component it's built on)
|
||||
import { Form } from "antd";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi, beforeEach } from "vitest";
|
||||
import { renderWithProviders, screen } from "../../../tests/test-utils";
|
||||
import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils";
|
||||
import KeyLifecycleSettings from "./KeyLifecycleSettings";
|
||||
|
||||
vi.mock("antd", () => {
|
||||
const Option = ({ children, value }: any) => <option value={value}>{children}</option>;
|
||||
const Select = ({ children, value, onChange, placeholder }: any) => (
|
||||
<select
|
||||
data-testid="select"
|
||||
value={value}
|
||||
onChange={(e) => onChange(e.target.value)}
|
||||
data-placeholder={placeholder}
|
||||
>
|
||||
{children}
|
||||
</select>
|
||||
);
|
||||
Select.Option = Option;
|
||||
return {
|
||||
Select,
|
||||
Tooltip: ({ children, title }: any) => (
|
||||
<div data-testid="tooltip" title={title}>
|
||||
{children}
|
||||
</div>
|
||||
),
|
||||
Switch: ({ checked, onChange }: any) => (
|
||||
<input type="checkbox" data-testid="switch" checked={checked} onChange={(e) => onChange(e.target.checked)} />
|
||||
),
|
||||
Divider: () => <hr data-testid="divider" />,
|
||||
};
|
||||
});
|
||||
const CREATE_PLACEHOLDER = "e.g., 30d or leave empty to never expire";
|
||||
const EDIT_PLACEHOLDER = "e.g., 30d";
|
||||
|
||||
vi.mock("@ant-design/icons", () => ({
|
||||
InfoCircleOutlined: () => <span data-testid="info-icon">ℹ</span>,
|
||||
}));
|
||||
interface HarnessProps {
|
||||
isCreateMode?: boolean;
|
||||
onFinish?: (values: Record<string, unknown>) => void;
|
||||
}
|
||||
|
||||
vi.mock("@tremor/react", () => ({
|
||||
TextInput: ({ value, onValueChange, onChange, placeholder, name, className }: any) => {
|
||||
const handleChange = (e: React.ChangeEvent<HTMLInputElement>) => {
|
||||
if (onChange) {
|
||||
onChange(e);
|
||||
}
|
||||
if (onValueChange) {
|
||||
onValueChange(e.target.value);
|
||||
}
|
||||
};
|
||||
return (
|
||||
<input
|
||||
data-testid={name === "duration" ? "duration-input" : "custom-interval-input"}
|
||||
value={value}
|
||||
onChange={handleChange}
|
||||
placeholder={placeholder}
|
||||
className={className}
|
||||
const Harness: React.FC<HarnessProps> = ({ isCreateMode = true, onFinish = () => {} }) => {
|
||||
const [form] = Form.useForm();
|
||||
const [autoRotationEnabled, setAutoRotationEnabled] = useState(false);
|
||||
const [rotationInterval, setRotationInterval] = useState("");
|
||||
const [neverExpire, setNeverExpire] = useState(false);
|
||||
|
||||
return (
|
||||
<Form form={form} onFinish={onFinish}>
|
||||
<KeyLifecycleSettings
|
||||
form={form}
|
||||
autoRotationEnabled={autoRotationEnabled}
|
||||
onAutoRotationChange={setAutoRotationEnabled}
|
||||
rotationInterval={rotationInterval}
|
||||
onRotationIntervalChange={setRotationInterval}
|
||||
isCreateMode={isCreateMode}
|
||||
neverExpire={neverExpire}
|
||||
onNeverExpireChange={setNeverExpire}
|
||||
/>
|
||||
);
|
||||
},
|
||||
}));
|
||||
<button type="submit">submit</button>
|
||||
<button type="button" onClick={() => form.resetFields()}>
|
||||
reset
|
||||
</button>
|
||||
<span data-testid="rotation-interval-value">{rotationInterval}</span>
|
||||
</Form>
|
||||
);
|
||||
};
|
||||
|
||||
const getDurationInput = (isCreateMode = true) =>
|
||||
screen.getByPlaceholderText(isCreateMode ? CREATE_PLACEHOLDER : EDIT_PLACEHOLDER) as HTMLInputElement;
|
||||
|
||||
describe("KeyLifecycleSettings", () => {
|
||||
const mockForm = {
|
||||
getFieldValue: vi.fn(),
|
||||
setFieldValue: vi.fn(),
|
||||
setFieldsValue: vi.fn(),
|
||||
};
|
||||
|
||||
const defaultProps = {
|
||||
form: mockForm,
|
||||
autoRotationEnabled: false,
|
||||
onAutoRotationChange: vi.fn(),
|
||||
rotationInterval: "",
|
||||
onRotationIntervalChange: vi.fn(),
|
||||
isCreateMode: false,
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockForm.getFieldValue.mockReturnValue("");
|
||||
});
|
||||
|
||||
it("should render without crashing", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
|
||||
|
||||
it("renders the expiry and auto-rotation sections", () => {
|
||||
renderWithProviders(<Harness />);
|
||||
expect(screen.getByText("Key Expiry Settings")).toBeInTheDocument();
|
||||
expect(screen.getByText("Auto-Rotation Settings")).toBeInTheDocument();
|
||||
expect(getDurationInput()).toBeInTheDocument();
|
||||
});
|
||||
|
||||
describe("Key Expiry Settings", () => {
|
||||
it("should render expiry input field", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
|
||||
it("uses the create-mode placeholder in create mode", () => {
|
||||
renderWithProviders(<Harness isCreateMode={true} />);
|
||||
expect(screen.getByPlaceholderText(CREATE_PLACEHOLDER)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
expect(screen.getByText("Expire Key")).toBeInTheDocument();
|
||||
expect(screen.getByTestId("duration-input")).toBeInTheDocument();
|
||||
});
|
||||
it("uses the edit-mode placeholder in edit mode", () => {
|
||||
renderWithProviders(<Harness isCreateMode={false} />);
|
||||
expect(screen.getByPlaceholderText(EDIT_PLACEHOLDER)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show correct placeholder in create mode", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={true} />);
|
||||
|
||||
const input = screen.getByTestId("duration-input");
|
||||
expect(input).toHaveAttribute("placeholder", "e.g., 30d or leave empty to never expire");
|
||||
});
|
||||
|
||||
it("should show correct placeholder in edit mode", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={false} />);
|
||||
|
||||
const input = screen.getByTestId("duration-input");
|
||||
expect(input).toHaveAttribute("placeholder", "e.g., 30d");
|
||||
});
|
||||
|
||||
it("should show correct tooltip in create mode", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={true} />);
|
||||
|
||||
const tooltips = screen.getAllByTestId("tooltip");
|
||||
const expiryTooltip = tooltips.find((tooltip) =>
|
||||
tooltip.getAttribute("title")?.includes("Leave empty to keep the current expiry unchanged"),
|
||||
);
|
||||
expect(expiryTooltip).toBeInTheDocument();
|
||||
expect(expiryTooltip).toHaveAttribute(
|
||||
"title",
|
||||
"Set when this key should expire. Format: 30s (seconds), 30m (minutes), 30h (hours), 30d (days). Leave empty to keep the current expiry unchanged.",
|
||||
);
|
||||
});
|
||||
|
||||
it("should show correct tooltip in edit mode", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={false} />);
|
||||
|
||||
const tooltips = screen.getAllByTestId("tooltip");
|
||||
const expiryTooltip = tooltips.find((tooltip) =>
|
||||
tooltip.getAttribute("title")?.includes("Leave empty to keep the current expiry unchanged"),
|
||||
);
|
||||
expect(expiryTooltip).toBeInTheDocument();
|
||||
expect(expiryTooltip).toHaveAttribute(
|
||||
"title",
|
||||
"Set when this key should expire. Format: 30s (seconds), 30m (minutes), 30h (hours), 30d (days). Leave empty to keep the current expiry unchanged.",
|
||||
);
|
||||
});
|
||||
|
||||
it("should initialize with form value if present", () => {
|
||||
mockForm.getFieldValue.mockReturnValue("30d");
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
|
||||
|
||||
const input = screen.getByTestId("duration-input") as HTMLInputElement;
|
||||
expect(input.value).toBe("30d");
|
||||
});
|
||||
|
||||
it("should update form using setFieldValue when duration changes", async () => {
|
||||
describe("duration is a single source of truth (regression for pre-filled value dropped on submit)", () => {
|
||||
it("submits the duration the user typed", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
|
||||
const onFinish = vi.fn();
|
||||
renderWithProviders(<Harness onFinish={onFinish} />);
|
||||
|
||||
const input = screen.getByTestId("duration-input");
|
||||
await user.type(input, "60d");
|
||||
await user.type(getDurationInput(), "1d");
|
||||
await user.click(screen.getByRole("button", { name: "submit" }));
|
||||
|
||||
expect(mockForm.setFieldValue).toHaveBeenCalledWith("duration", "60d");
|
||||
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(1));
|
||||
expect(onFinish.mock.calls[0][0]).toMatchObject({ duration: "1d" });
|
||||
});
|
||||
|
||||
it("should update form using setFieldsValue when setFieldValue is not available", async () => {
|
||||
it("clears the displayed value when the form is reset, so no stale value lingers", async () => {
|
||||
const user = userEvent.setup();
|
||||
const formWithoutSetFieldValue = {
|
||||
getFieldValue: vi.fn().mockReturnValue(""),
|
||||
setFieldsValue: vi.fn(),
|
||||
};
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} form={formWithoutSetFieldValue} />);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
const input = screen.getByTestId("duration-input");
|
||||
await user.type(input, "90d");
|
||||
await user.type(getDurationInput(), "1d");
|
||||
expect(getDurationInput().value).toBe("1d");
|
||||
|
||||
expect(formWithoutSetFieldValue.setFieldsValue).toHaveBeenCalledWith({ duration: "90d" });
|
||||
await user.click(screen.getByRole("button", { name: "reset" }));
|
||||
|
||||
await waitFor(() => expect(getDurationInput().value).toBe(""));
|
||||
});
|
||||
|
||||
it("never submits a value that differs from what is displayed after a reset", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onFinish = vi.fn();
|
||||
renderWithProviders(<Harness onFinish={onFinish} />);
|
||||
|
||||
// First create: type "1d" and submit -> "1d" is sent.
|
||||
await user.type(getDurationInput(), "1d");
|
||||
await user.click(screen.getByRole("button", { name: "submit" }));
|
||||
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(1));
|
||||
expect(onFinish.mock.calls[0][0]).toMatchObject({ duration: "1d" });
|
||||
|
||||
// Second create: form resets, so the field must show empty AND submit empty.
|
||||
// The old bug showed a stale "1d" while submitting null/empty.
|
||||
await user.click(screen.getByRole("button", { name: "reset" }));
|
||||
await waitFor(() => expect(getDurationInput().value).toBe(""));
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "submit" }));
|
||||
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(2));
|
||||
expect(onFinish.mock.calls[1][0].duration).not.toBe("1d");
|
||||
expect(getDurationInput().value).toBe(onFinish.mock.calls[1][0].duration ?? "");
|
||||
});
|
||||
});
|
||||
|
||||
describe("Auto-Rotation Settings", () => {
|
||||
it("should render auto-rotation switch", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
|
||||
|
||||
expect(screen.getByText("Enable Auto-Rotation")).toBeInTheDocument();
|
||||
expect(screen.getByTestId("switch")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show switch as unchecked when autoRotationEnabled is false", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={false} />);
|
||||
|
||||
const switchElement = screen.getByTestId("switch") as HTMLInputElement;
|
||||
expect(switchElement.checked).toBe(false);
|
||||
});
|
||||
|
||||
it("should show switch as checked when autoRotationEnabled is true", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} />);
|
||||
|
||||
const switchElement = screen.getByTestId("switch") as HTMLInputElement;
|
||||
expect(switchElement.checked).toBe(true);
|
||||
});
|
||||
|
||||
it("should call onAutoRotationChange when switch is toggled", async () => {
|
||||
describe("Never Expire", () => {
|
||||
it("clears and disables the duration input, then submits an empty duration", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onAutoRotationChange = vi.fn();
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} onAutoRotationChange={onAutoRotationChange} />);
|
||||
const onFinish = vi.fn();
|
||||
renderWithProviders(<Harness isCreateMode={false} onFinish={onFinish} />);
|
||||
|
||||
const switchElement = screen.getByTestId("switch");
|
||||
await user.click(switchElement);
|
||||
await user.type(getDurationInput(false), "30d");
|
||||
expect(getDurationInput(false).value).toBe("30d");
|
||||
|
||||
expect(onAutoRotationChange).toHaveBeenCalledWith(true);
|
||||
await user.click(screen.getByRole("checkbox", { name: /never expire/i }));
|
||||
|
||||
await waitFor(() => expect(getDurationInput(false).value).toBe(""));
|
||||
expect(getDurationInput(false)).toBeDisabled();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "submit" }));
|
||||
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(1));
|
||||
expect(onFinish.mock.calls[0][0]).toMatchObject({ duration: "" });
|
||||
});
|
||||
});
|
||||
|
||||
it("should not show rotation interval section when auto-rotation is disabled", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={false} />);
|
||||
describe("Auto-Rotation", () => {
|
||||
it("reveals the rotation interval controls when enabled", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
expect(screen.queryByText("Rotation Interval")).not.toBeInTheDocument();
|
||||
expect(screen.queryByTestId("select")).not.toBeInTheDocument();
|
||||
await user.click(screen.getByRole("switch"));
|
||||
|
||||
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
|
||||
});
|
||||
|
||||
it("should show rotation interval section when auto-rotation is enabled", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="30d" />);
|
||||
|
||||
expect(screen.getByText("Rotation Interval")).toBeInTheDocument();
|
||||
expect(screen.getByTestId("select")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show all predefined interval options", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="30d" />);
|
||||
|
||||
expect(screen.getByText("7 days")).toBeInTheDocument();
|
||||
expect(screen.getByText("30 days")).toBeInTheDocument();
|
||||
expect(screen.getByText("90 days")).toBeInTheDocument();
|
||||
expect(screen.getByText("180 days")).toBeInTheDocument();
|
||||
expect(screen.getByText("365 days")).toBeInTheDocument();
|
||||
expect(screen.getByText("Custom interval")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display current rotation interval in select", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="90d" />);
|
||||
|
||||
const select = screen.getByTestId("select") as HTMLSelectElement;
|
||||
expect(select.value).toBe("90d");
|
||||
});
|
||||
|
||||
it("should call onRotationIntervalChange when predefined interval is selected", async () => {
|
||||
it("propagates a selected predefined interval", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onRotationIntervalChange = vi.fn();
|
||||
renderWithProviders(
|
||||
<KeyLifecycleSettings
|
||||
{...defaultProps}
|
||||
autoRotationEnabled={true}
|
||||
rotationInterval="7d"
|
||||
onRotationIntervalChange={onRotationIntervalChange}
|
||||
/>,
|
||||
);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
const select = screen.getByTestId("select");
|
||||
await user.selectOptions(select, "30d");
|
||||
await user.click(screen.getByRole("switch"));
|
||||
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
|
||||
|
||||
expect(onRotationIntervalChange).toHaveBeenCalledWith("30d");
|
||||
await user.click(screen.getByRole("combobox"));
|
||||
await user.click(await screen.findByText("90 days"));
|
||||
|
||||
await waitFor(() => expect(document.querySelector(".ant-select-selection-item")?.textContent).toBe("90 days"));
|
||||
expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("90d");
|
||||
});
|
||||
|
||||
it("should show custom input when custom option is selected", async () => {
|
||||
it("shows the custom interval input when Custom interval is selected, without propagating yet", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="30d" />);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
const select = screen.getByTestId("select");
|
||||
await user.selectOptions(select, "custom");
|
||||
await user.click(screen.getByRole("switch"));
|
||||
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
|
||||
|
||||
expect(screen.getByTestId("custom-interval-input")).toBeInTheDocument();
|
||||
await user.click(screen.getByRole("combobox"));
|
||||
await user.click(await screen.findByText("Custom interval"));
|
||||
|
||||
expect(await screen.findByPlaceholderText("e.g., 1s, 5m, 2h, 14d")).toBeInTheDocument();
|
||||
expect(screen.getByText("Supported formats: seconds (s), minutes (m), hours (h), days (d)")).toBeInTheDocument();
|
||||
expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("");
|
||||
});
|
||||
|
||||
it("should hide custom input when predefined interval is selected after custom", async () => {
|
||||
it("propagates a typed custom interval to the parent", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onRotationIntervalChange = vi.fn();
|
||||
renderWithProviders(
|
||||
<KeyLifecycleSettings
|
||||
{...defaultProps}
|
||||
autoRotationEnabled={true}
|
||||
rotationInterval="custom-value"
|
||||
onRotationIntervalChange={onRotationIntervalChange}
|
||||
/>,
|
||||
);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
const select = screen.getByTestId("select");
|
||||
await user.selectOptions(select, "7d");
|
||||
await user.click(screen.getByRole("switch"));
|
||||
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
|
||||
|
||||
expect(screen.queryByTestId("custom-interval-input")).not.toBeInTheDocument();
|
||||
expect(onRotationIntervalChange).toHaveBeenCalledWith("7d");
|
||||
});
|
||||
await user.click(screen.getByRole("combobox"));
|
||||
await user.click(await screen.findByText("Custom interval"));
|
||||
|
||||
it("should call onRotationIntervalChange when custom interval is entered", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onRotationIntervalChange = vi.fn();
|
||||
renderWithProviders(
|
||||
<KeyLifecycleSettings
|
||||
{...defaultProps}
|
||||
autoRotationEnabled={true}
|
||||
rotationInterval=""
|
||||
onRotationIntervalChange={onRotationIntervalChange}
|
||||
/>,
|
||||
);
|
||||
|
||||
const select = screen.getByTestId("select");
|
||||
await user.selectOptions(select, "custom");
|
||||
|
||||
const customInput = screen.getByTestId("custom-interval-input");
|
||||
const customInput = await screen.findByPlaceholderText("e.g., 1s, 5m, 2h, 14d");
|
||||
await user.type(customInput, "14d");
|
||||
|
||||
expect(onRotationIntervalChange).toHaveBeenCalledWith("14d");
|
||||
await waitFor(() => expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("14d"));
|
||||
expect((customInput as HTMLInputElement).value).toBe("14d");
|
||||
});
|
||||
|
||||
it("should show info message when auto-rotation is enabled", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} />);
|
||||
|
||||
expect(
|
||||
screen.getByText(
|
||||
"When rotation occurs, you'll receive a notification with the new key. The old key will be deactivated after a brief grace period.",
|
||||
),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show info message when auto-rotation is disabled", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={false} />);
|
||||
|
||||
expect(
|
||||
screen.queryByText(
|
||||
"When rotation occurs, you'll receive a notification with the new key. The old key will be deactivated after a brief grace period.",
|
||||
),
|
||||
).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should initialize with custom interval input visible when custom interval is provided", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="14d" />);
|
||||
|
||||
expect(screen.getByTestId("custom-interval-input")).toBeInTheDocument();
|
||||
const customInput = screen.getByTestId("custom-interval-input") as HTMLInputElement;
|
||||
expect(customInput.value).toBe("14d");
|
||||
});
|
||||
|
||||
it("should show custom option selected when custom interval is provided", () => {
|
||||
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="14d" />);
|
||||
|
||||
const select = screen.getByTestId("select") as HTMLSelectElement;
|
||||
expect(select.value).toBe("custom");
|
||||
});
|
||||
|
||||
it("should not call onRotationIntervalChange when selecting custom option", async () => {
|
||||
it("hides the custom input and propagates the value when switching back to a predefined interval", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onRotationIntervalChange = vi.fn();
|
||||
renderWithProviders(
|
||||
<KeyLifecycleSettings
|
||||
{...defaultProps}
|
||||
autoRotationEnabled={true}
|
||||
rotationInterval="30d"
|
||||
onRotationIntervalChange={onRotationIntervalChange}
|
||||
/>,
|
||||
);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
const select = screen.getByTestId("select");
|
||||
await user.selectOptions(select, "custom");
|
||||
await user.click(screen.getByRole("switch"));
|
||||
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
|
||||
|
||||
expect(onRotationIntervalChange).not.toHaveBeenCalled();
|
||||
await user.click(screen.getByRole("combobox"));
|
||||
await user.click(await screen.findByText("Custom interval"));
|
||||
const customInput = await screen.findByPlaceholderText("e.g., 1s, 5m, 2h, 14d");
|
||||
await user.type(customInput, "14d");
|
||||
await waitFor(() => expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("14d"));
|
||||
|
||||
await user.click(screen.getByRole("combobox"));
|
||||
await user.click(await screen.findByText("7 days"));
|
||||
|
||||
await waitFor(() => expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("7d"));
|
||||
expect(screen.queryByPlaceholderText("e.g., 1s, 5m, 2h, 14d")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import React, { useState } from "react";
|
||||
import { Select, Tooltip, Divider, Switch, Checkbox } from "antd";
|
||||
import { Select, Tooltip, Divider, Switch, Checkbox, Form } from "antd";
|
||||
import { InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { TextInput } from "@tremor/react";
|
||||
|
||||
|
|
@ -34,7 +34,6 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
|
|||
|
||||
const [showCustomInput, setShowCustomInput] = useState(isCustomInterval);
|
||||
const [customInterval, setCustomInterval] = useState(isCustomInterval ? rotationInterval : "");
|
||||
const [durationValue, setDurationValue] = useState<string>(form?.getFieldValue?.("duration") || "");
|
||||
|
||||
const handleIntervalChange = (value: string) => {
|
||||
if (value === "custom") {
|
||||
|
|
@ -53,14 +52,6 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
|
|||
onRotationIntervalChange(value);
|
||||
};
|
||||
|
||||
const handleDurationChange = (value: string) => {
|
||||
setDurationValue(value);
|
||||
if (form && typeof form.setFieldValue === "function") {
|
||||
form.setFieldValue("duration", value);
|
||||
} else if (form && typeof form.setFieldsValue === "function") {
|
||||
form.setFieldsValue({ duration: value });
|
||||
}
|
||||
};
|
||||
return (
|
||||
<div className="space-y-6">
|
||||
{/* Key Expiry Section */}
|
||||
|
|
@ -80,7 +71,6 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
|
|||
const checked = e.target.checked;
|
||||
onNeverExpireChange(checked);
|
||||
if (checked) {
|
||||
setDurationValue("");
|
||||
if (form && typeof form.setFieldValue === "function") {
|
||||
form.setFieldValue("duration", "");
|
||||
} else if (form && typeof form.setFieldsValue === "function") {
|
||||
|
|
@ -94,14 +84,13 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
|
|||
</Checkbox>
|
||||
)}
|
||||
</label>
|
||||
<TextInput
|
||||
name="duration"
|
||||
placeholder={isCreateMode ? "e.g., 30d or leave empty to never expire" : "e.g., 30d"}
|
||||
className="w-full"
|
||||
value={durationValue}
|
||||
onValueChange={handleDurationChange}
|
||||
disabled={!isCreateMode && neverExpire}
|
||||
/>
|
||||
<Form.Item name="duration" noStyle initialValue="">
|
||||
<TextInput
|
||||
placeholder={isCreateMode ? "e.g., 30d or leave empty to never expire" : "e.g., 30d"}
|
||||
className="w-full"
|
||||
disabled={!isCreateMode && neverExpire}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import { Setter } from "@/types";
|
||||
import { useEffect, useState } from "react";
|
||||
import { keyListCall, Member, Organization } from "../networking";
|
||||
import type { ObjectPermission } from "../object_permission_types";
|
||||
|
||||
export interface Team {
|
||||
team_id: string;
|
||||
|
|
@ -90,15 +91,7 @@ export interface KeyResponse {
|
|||
user_tpm_limit: number;
|
||||
user_rpm_limit: number;
|
||||
user_email: string;
|
||||
object_permission?: {
|
||||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
vector_stores: string[];
|
||||
agents?: string[];
|
||||
agent_access_groups?: string[];
|
||||
};
|
||||
object_permission?: ObjectPermission | null;
|
||||
access_group_ids?: string[];
|
||||
budget_fallbacks?: Record<string, string[]>;
|
||||
budget_limits?: Array<{ budget_duration: string; max_budget: number; reset_at?: string }>;
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ import { TagNewRequest, TagUpdateRequest, TagListResponse, TagInfoResponse } fro
|
|||
import { Team } from "./key_team_helpers/key_list";
|
||||
import { EmailEventSettingsResponse, EmailEventSettingsUpdateRequest } from "./email_events/types";
|
||||
import type { SkillRegisterRequest } from "./claude_code_plugins/types";
|
||||
import type { ObjectPermission } from "./object_permission_types";
|
||||
import { jsonFields } from "./common_components/check_openapi_schema";
|
||||
import NotificationsManager from "./molecules/notifications_manager";
|
||||
import type { MCPUserEnvVarsStatus } from "./mcp_tools/types";
|
||||
|
|
@ -208,13 +209,7 @@ export interface Organization {
|
|||
teams: any[] | null;
|
||||
users: any[] | null;
|
||||
members: any[] | null;
|
||||
object_permission?: {
|
||||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_toolsets?: string[];
|
||||
vector_stores: string[];
|
||||
};
|
||||
object_permission?: ObjectPermission | null;
|
||||
}
|
||||
|
||||
export interface CredentialItem {
|
||||
|
|
@ -1184,35 +1179,6 @@ export const organizationInfoCall = async (accessToken: string, organizationID:
|
|||
}
|
||||
};
|
||||
|
||||
export const organizationCreateCall = async (
|
||||
accessToken: string,
|
||||
formValues: Record<string, any>, // Assuming formValues is an object
|
||||
) => {
|
||||
try {
|
||||
if (formValues.metadata) {
|
||||
// if there's an exception JSON.parse, show it in the message
|
||||
try {
|
||||
formValues.metadata = JSON.parse(formValues.metadata);
|
||||
} catch (error) {
|
||||
console.error("Failed to parse metadata:", error);
|
||||
throw new Error("Failed to parse metadata: " + error);
|
||||
}
|
||||
}
|
||||
|
||||
const data = await apiClient.post(`/organization/new`, {
|
||||
accessToken,
|
||||
body: {
|
||||
...formValues, // Include formValues in the request body
|
||||
},
|
||||
});
|
||||
return data;
|
||||
// Handle success - you might want to update some state or UI based on the created key
|
||||
} catch (error) {
|
||||
console.error("Failed to create key:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const organizationUpdateCall = async (
|
||||
accessToken: string,
|
||||
formValues: Record<string, any>, // Assuming formValues is an object
|
||||
|
|
|
|||
|
|
@ -0,0 +1,3 @@
|
|||
import type { components } from "@/lib/http/schema";
|
||||
|
||||
export type ObjectPermission = Partial<components["schemas"]["LiteLLM_ObjectPermissionTable"]>;
|
||||
|
|
@ -3,21 +3,10 @@ import { Text } from "@tremor/react";
|
|||
import VectorStorePermissions from "./permissions/VectorStorePermissions";
|
||||
import MCPServerPermissions from "./permissions/MCPServerPermissions";
|
||||
import AgentPermissions from "./permissions/AgentPermissions";
|
||||
|
||||
interface ObjectPermission {
|
||||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
mcp_toolsets?: string[];
|
||||
vector_stores: string[];
|
||||
agents?: string[];
|
||||
agent_access_groups?: string[];
|
||||
search_tools?: string[];
|
||||
}
|
||||
import type { ObjectPermission } from "./object_permission_types";
|
||||
|
||||
interface ObjectPermissionsViewProps {
|
||||
objectPermission?: ObjectPermission;
|
||||
objectPermission?: ObjectPermission | null;
|
||||
variant?: "card" | "inline";
|
||||
className?: string;
|
||||
accessToken?: string | null;
|
||||
|
|
|
|||
|
|
@ -443,6 +443,36 @@ describe("CreateKey", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("should include mcp_toolsets in keyCreateCall payload when only toolsets are selected", async () => {
|
||||
renderWithProviders(<CreateKey {...defaultProps} />);
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByRole("button", { name: /create new key/i }));
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("button", { name: /create key/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
act(() => {
|
||||
formMock.setFieldValue("key_alias", "Test Key");
|
||||
formMock.setFieldValue("allowed_mcp_servers_and_groups", {
|
||||
servers: [],
|
||||
accessGroups: [],
|
||||
toolsets: ["ts-1"],
|
||||
});
|
||||
});
|
||||
|
||||
act(() => {
|
||||
fireEvent.click(screen.getByRole("button", { name: /create key/i }));
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockKeyCreateCall).toHaveBeenCalled();
|
||||
});
|
||||
expect(mockKeyCreateCall.mock.calls[0][2].object_permission?.mcp_toolsets).toEqual(["ts-1"]);
|
||||
});
|
||||
|
||||
it("should prefill models when provided without team_id", async () => {
|
||||
renderWithProviders(
|
||||
<CreateKey
|
||||
|
|
|
|||
|
|
@ -468,18 +468,22 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
if (
|
||||
formValues.allowed_mcp_servers_and_groups &&
|
||||
(formValues.allowed_mcp_servers_and_groups.servers?.length > 0 ||
|
||||
formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0)
|
||||
formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0 ||
|
||||
formValues.allowed_mcp_servers_and_groups.toolsets?.length > 0)
|
||||
) {
|
||||
if (!formValues.object_permission) {
|
||||
formValues.object_permission = {};
|
||||
}
|
||||
const { servers, accessGroups } = formValues.allowed_mcp_servers_and_groups;
|
||||
const { servers, accessGroups, toolsets } = formValues.allowed_mcp_servers_and_groups;
|
||||
if (servers && servers.length > 0) {
|
||||
formValues.object_permission.mcp_servers = servers;
|
||||
}
|
||||
if (accessGroups && accessGroups.length > 0) {
|
||||
formValues.object_permission.mcp_access_groups = accessGroups;
|
||||
}
|
||||
if (toolsets && toolsets.length > 0) {
|
||||
formValues.object_permission.mcp_toolsets = toolsets;
|
||||
}
|
||||
// Remove the original field as it's now part of object_permission
|
||||
delete formValues.allowed_mcp_servers_and_groups;
|
||||
}
|
||||
|
|
@ -1640,9 +1644,6 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
|
|||
/>
|
||||
</div>
|
||||
</AccordionBody>
|
||||
<Form.Item name="duration" hidden initialValue={null}>
|
||||
<Input />
|
||||
</Form.Item>
|
||||
</Accordion>
|
||||
<Accordion className="mt-4 mb-4">
|
||||
<AccordionHeader>
|
||||
|
|
|
|||
|
|
@ -0,0 +1,203 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import React from "react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
vi.mock("@/components/molecules/notifications_manager", () => ({
|
||||
__esModule: true,
|
||||
default: { success: vi.fn(), fromBackend: vi.fn() },
|
||||
}));
|
||||
vi.mock("@/components/ModelSelect/ModelSelect", () => ({
|
||||
ModelSelect: ({ onChange }: { onChange: (values: string[]) => void }) => (
|
||||
<button type="button" onClick={() => onChange(["gpt-5.2"])}>
|
||||
set-models
|
||||
</button>
|
||||
),
|
||||
}));
|
||||
vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({
|
||||
__esModule: true,
|
||||
default: ({ onChange }: { onChange: (values: string[]) => void }) => (
|
||||
<button type="button" onClick={() => onChange(["vs-1"])}>
|
||||
set-vector-stores
|
||||
</button>
|
||||
),
|
||||
}));
|
||||
vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({
|
||||
__esModule: true,
|
||||
default: ({
|
||||
onChange,
|
||||
}: {
|
||||
onChange: (values: { servers: string[]; accessGroups: string[]; toolsets: string[] }) => void;
|
||||
}) => (
|
||||
<button type="button" onClick={() => onChange({ servers: ["srv-1"], accessGroups: [], toolsets: ["ts-1"] })}>
|
||||
set-mcp
|
||||
</button>
|
||||
),
|
||||
}));
|
||||
|
||||
import { OrgCreateDialog } from "./OrgCreateDialog";
|
||||
|
||||
const Harness = ({ createOrganization }: { createOrganization: (body: unknown) => Promise<unknown> }) => {
|
||||
const [open, setOpen] = React.useState(true);
|
||||
return (
|
||||
<>
|
||||
<button type="button" onClick={() => setOpen(true)}>
|
||||
reopen
|
||||
</button>
|
||||
<OrgCreateDialog open={open} onOpenChange={setOpen} accessToken="token" createOrganization={createOrganization} />
|
||||
</>
|
||||
);
|
||||
};
|
||||
|
||||
const renderDialog = (overrides?: { createOrganization?: ReturnType<typeof vi.fn> }) => {
|
||||
const createOrganization = overrides?.createOrganization ?? vi.fn().mockResolvedValue({});
|
||||
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||
render(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<Harness createOrganization={createOrganization} />
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
return { createOrganization };
|
||||
};
|
||||
|
||||
describe("OrgCreateDialog", () => {
|
||||
it("blocks submit and shows an error when the name is missing", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createOrganization } = renderDialog();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "Create Organization" }));
|
||||
|
||||
expect(await screen.findByRole("alert")).toHaveTextContent("Please input an organization name");
|
||||
expect(createOrganization).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("sends only alias and models for a minimal create and closes the dialog", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createOrganization } = renderDialog();
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "new-org");
|
||||
await user.click(screen.getByRole("button", { name: "Create Organization" }));
|
||||
|
||||
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
|
||||
expect(createOrganization.mock.calls[0][0]).toStrictEqual({ organization_alias: "new-org", models: [] });
|
||||
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
|
||||
});
|
||||
|
||||
it("maps selectors and limits into the create body", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createOrganization } = renderDialog();
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "new-org");
|
||||
await user.click(screen.getByRole("button", { name: "set-models" }));
|
||||
await user.type(screen.getByLabelText("Tokens per minute Limit (TPM)"), "1000");
|
||||
await user.click(screen.getByRole("button", { name: "set-vector-stores" }));
|
||||
await user.click(screen.getByRole("button", { name: "set-mcp" }));
|
||||
await user.click(screen.getByRole("button", { name: "Create Organization" }));
|
||||
|
||||
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
|
||||
const expectedBody = {
|
||||
organization_alias: "new-org",
|
||||
models: ["gpt-5.2"],
|
||||
tpm_limit: 1000,
|
||||
object_permission: {
|
||||
vector_stores: ["vs-1"],
|
||||
mcp_servers: ["srv-1"],
|
||||
mcp_toolsets: ["ts-1"],
|
||||
},
|
||||
};
|
||||
expect(createOrganization.mock.calls[0][0]).toStrictEqual(expectedBody);
|
||||
});
|
||||
|
||||
it("blocks submit and shows an error for invalid metadata JSON", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createOrganization } = renderDialog();
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "new-org");
|
||||
await user.type(screen.getByLabelText("Metadata"), "not json");
|
||||
await user.click(screen.getByRole("button", { name: "Create Organization" }));
|
||||
|
||||
expect(await screen.findByRole("alert")).toHaveTextContent("Metadata must be a valid JSON object");
|
||||
expect(createOrganization).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("keeps the dialog open with the entered values when the create fails", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createOrganization } = renderDialog({
|
||||
createOrganization: vi.fn().mockRejectedValue(new Error("boom")),
|
||||
});
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "new-org");
|
||||
await user.click(screen.getByRole("button", { name: "Create Organization" }));
|
||||
|
||||
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
|
||||
expect(screen.getByLabelText("Organization Name")).toHaveValue("new-org");
|
||||
});
|
||||
|
||||
it("resets the form when the dialog is cancelled and reopened", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderDialog();
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "abandoned");
|
||||
await user.click(screen.getByRole("button", { name: "Cancel" }));
|
||||
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "reopen" }));
|
||||
expect(screen.getByLabelText("Organization Name")).toHaveValue("");
|
||||
});
|
||||
|
||||
it("resets the form when the dialog is dismissed with Escape and reopened", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderDialog();
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "abandoned");
|
||||
await user.keyboard("{Escape}");
|
||||
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "reopen" }));
|
||||
expect(screen.getByLabelText("Organization Name")).toHaveValue("");
|
||||
});
|
||||
|
||||
it("cannot be dismissed while a create is pending, then closes once on success", async () => {
|
||||
const user = userEvent.setup();
|
||||
let resolveCreate: (value: unknown) => void = () => {};
|
||||
const createOrganization = vi.fn().mockImplementation(
|
||||
() =>
|
||||
new Promise((resolve) => {
|
||||
resolveCreate = resolve;
|
||||
}),
|
||||
);
|
||||
renderDialog({ createOrganization });
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "new-org");
|
||||
await user.keyboard("{Enter}");
|
||||
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
|
||||
|
||||
await user.keyboard("{Escape}");
|
||||
expect(screen.getByLabelText("Organization Name")).toHaveValue("new-org");
|
||||
|
||||
resolveCreate({});
|
||||
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
|
||||
});
|
||||
|
||||
it("does not fire a second create while one is pending", async () => {
|
||||
const user = userEvent.setup();
|
||||
let resolveCreate: (value: unknown) => void = () => {};
|
||||
const createOrganization = vi.fn().mockImplementation(
|
||||
() =>
|
||||
new Promise((resolve) => {
|
||||
resolveCreate = resolve;
|
||||
}),
|
||||
);
|
||||
renderDialog({ createOrganization });
|
||||
|
||||
await user.type(screen.getByLabelText("Organization Name"), "new-org");
|
||||
await user.keyboard("{Enter}");
|
||||
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
|
||||
await user.keyboard("{Enter}");
|
||||
|
||||
expect(createOrganization).toHaveBeenCalledTimes(1);
|
||||
resolveCreate({});
|
||||
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,186 @@
|
|||
"use client";
|
||||
|
||||
import { useMutation, useQueryClient } from "@tanstack/react-query";
|
||||
import * as React from "react";
|
||||
|
||||
import { organizationKeys } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
|
||||
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
|
||||
import NotificationsManager from "@/components/molecules/notifications_manager";
|
||||
import { FieldGroup } from "@/components/shared/form/field";
|
||||
import { FormField } from "@/components/shared/form/FormField";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector";
|
||||
import { useZodForm } from "@/lib/forms/useZodForm";
|
||||
import { fetchClient } from "@/lib/http/api";
|
||||
|
||||
import { BUDGET_DURATION_OPTIONS, NO_RESET } from "../org-settings/OrgSettingsForm";
|
||||
import { orgSettingsSchema } from "../org-settings/schema";
|
||||
import { buildOrgCreateBody, emptyOrgFormValues, type OrgCreateBody } from "./mapper";
|
||||
|
||||
const defaultCreateOrganization = async (body: OrgCreateBody): Promise<unknown> => {
|
||||
const { data } = await fetchClient.POST("/organization/new", { body });
|
||||
return data;
|
||||
};
|
||||
|
||||
interface OrgCreateDialogProps {
|
||||
open: boolean;
|
||||
onOpenChange: (open: boolean) => void;
|
||||
accessToken: string;
|
||||
createOrganization?: (body: OrgCreateBody) => Promise<unknown>;
|
||||
}
|
||||
|
||||
export const OrgCreateDialog = ({
|
||||
open,
|
||||
onOpenChange,
|
||||
accessToken,
|
||||
createOrganization = defaultCreateOrganization,
|
||||
}: OrgCreateDialogProps) => {
|
||||
const queryClient = useQueryClient();
|
||||
const form = useZodForm(orgSettingsSchema, { defaultValues: emptyOrgFormValues });
|
||||
|
||||
const closeAndReset = () => {
|
||||
form.reset(emptyOrgFormValues);
|
||||
onOpenChange(false);
|
||||
};
|
||||
|
||||
const mutation = useMutation({
|
||||
mutationFn: (body: OrgCreateBody) => createOrganization(body),
|
||||
onSuccess: () => {
|
||||
NotificationsManager.success("Organization created successfully");
|
||||
queryClient.invalidateQueries({ queryKey: organizationKeys.all });
|
||||
closeAndReset();
|
||||
},
|
||||
onError: (error: unknown) =>
|
||||
NotificationsManager.fromBackend(error instanceof Error ? error.message : "Failed to create organization"),
|
||||
});
|
||||
|
||||
const handleOpenChange = (nextOpen: boolean) => {
|
||||
if (!nextOpen && mutation.isPending) return;
|
||||
if (!nextOpen) {
|
||||
form.reset(emptyOrgFormValues);
|
||||
}
|
||||
onOpenChange(nextOpen);
|
||||
};
|
||||
|
||||
const onSubmit = form.handleSubmit((values) => {
|
||||
if (mutation.isPending) return;
|
||||
mutation.mutate(buildOrgCreateBody(values));
|
||||
});
|
||||
|
||||
return (
|
||||
<Dialog open={open} onOpenChange={handleOpenChange}>
|
||||
<DialogContent className="sm:max-w-3xl max-h-[90vh] overflow-y-auto">
|
||||
<DialogHeader>
|
||||
<DialogTitle>Create Organization</DialogTitle>
|
||||
</DialogHeader>
|
||||
|
||||
<form onSubmit={onSubmit}>
|
||||
<FieldGroup>
|
||||
<FormField control={form.control} name="organization_alias" label="Organization Name">
|
||||
{({ ref, ...field }) => <Input {...field} ref={ref} />}
|
||||
</FormField>
|
||||
|
||||
<FormField control={form.control} name="models" label="Models">
|
||||
{(field) => (
|
||||
<ModelSelect
|
||||
value={field.value}
|
||||
onChange={field.onChange}
|
||||
context="organization"
|
||||
options={{ includeSpecialOptions: true, showAllProxyModelsOverride: true }}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
||||
<FormField control={form.control} name="max_budget" label="Max Budget (USD)">
|
||||
{({ ref, ...field }) => <Input {...field} ref={ref} type="number" step={0.01} min={0} />}
|
||||
</FormField>
|
||||
|
||||
<FormField control={form.control} name="budget_duration" label="Reset Budget">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
<Select
|
||||
items={BUDGET_DURATION_OPTIONS}
|
||||
value={value === "" ? NO_RESET : value}
|
||||
onValueChange={(selected) => onChange(selected === NO_RESET ? "" : selected)}
|
||||
>
|
||||
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy}>
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{BUDGET_DURATION_OPTIONS.map((option) => (
|
||||
<SelectItem key={option.value} value={option.value}>
|
||||
{option.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
)}
|
||||
</FormField>
|
||||
|
||||
<FormField control={form.control} name="tpm_limit" label="Tokens per minute Limit (TPM)">
|
||||
{({ ref, ...field }) => <Input {...field} ref={ref} type="number" step={1} min={0} />}
|
||||
</FormField>
|
||||
|
||||
<FormField control={form.control} name="rpm_limit" label="Requests per minute Limit (RPM)">
|
||||
{({ ref, ...field }) => <Input {...field} ref={ref} type="number" step={1} min={0} />}
|
||||
</FormField>
|
||||
|
||||
<FormField
|
||||
control={form.control}
|
||||
name="vector_stores"
|
||||
label="Allowed Vector Stores"
|
||||
description="Select vector stores this organization can access. Leave empty for access to all vector stores"
|
||||
>
|
||||
{(field) => (
|
||||
<VectorStoreSelector
|
||||
value={field.value}
|
||||
onChange={field.onChange}
|
||||
accessToken={accessToken}
|
||||
placeholder="Select vector stores (optional)"
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
||||
<FormField
|
||||
control={form.control}
|
||||
name="mcp"
|
||||
label="Allowed MCP Servers"
|
||||
description="Select MCP servers, access groups, and toolsets this organization can access. Leave empty for access to all"
|
||||
>
|
||||
{(field) => (
|
||||
<MCPServerSelector
|
||||
value={field.value}
|
||||
onChange={field.onChange}
|
||||
accessToken={accessToken}
|
||||
placeholder="Select MCP servers and access groups (optional)"
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
||||
<FormField control={form.control} name="metadata" label="Metadata">
|
||||
{({ ref, ...field }) => <Textarea {...field} ref={ref} rows={4} />}
|
||||
</FormField>
|
||||
</FieldGroup>
|
||||
|
||||
<DialogFooter className="mt-6">
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
onClick={() => handleOpenChange(false)}
|
||||
disabled={mutation.isPending}
|
||||
>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button type="submit" disabled={mutation.isPending}>
|
||||
{mutation.isPending ? "Creating..." : "Create Organization"}
|
||||
</Button>
|
||||
</DialogFooter>
|
||||
</form>
|
||||
</DialogContent>
|
||||
</Dialog>
|
||||
);
|
||||
};
|
||||
|
|
@ -0,0 +1,59 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { buildOrgCreateBody, emptyOrgFormValues } from "./mapper";
|
||||
|
||||
describe("buildOrgCreateBody", () => {
|
||||
it("sends only alias and models for a minimal form", () => {
|
||||
expect(buildOrgCreateBody({ ...emptyOrgFormValues, organization_alias: "acme" })).toStrictEqual({
|
||||
organization_alias: "acme",
|
||||
models: [],
|
||||
});
|
||||
});
|
||||
|
||||
it("maps every field when the whole form is filled", () => {
|
||||
const filledForm = {
|
||||
organization_alias: "acme",
|
||||
models: ["gpt-5.2"],
|
||||
max_budget: "12.5",
|
||||
budget_duration: "30d",
|
||||
tpm_limit: "1000",
|
||||
rpm_limit: "50",
|
||||
vector_stores: ["vs-1"],
|
||||
mcp: { servers: ["srv-1"], accessGroups: ["ag-1"], toolsets: ["ts-1"] },
|
||||
metadata: '{"env": "prod"}',
|
||||
};
|
||||
const expectedBody = {
|
||||
organization_alias: "acme",
|
||||
models: ["gpt-5.2"],
|
||||
max_budget: 12.5,
|
||||
budget_duration: "30d",
|
||||
tpm_limit: 1000,
|
||||
rpm_limit: 50,
|
||||
metadata: { env: "prod" },
|
||||
object_permission: {
|
||||
vector_stores: ["vs-1"],
|
||||
mcp_servers: ["srv-1"],
|
||||
mcp_access_groups: ["ag-1"],
|
||||
mcp_toolsets: ["ts-1"],
|
||||
},
|
||||
};
|
||||
|
||||
expect(buildOrgCreateBody(filledForm)).toStrictEqual(expectedBody);
|
||||
});
|
||||
|
||||
it("includes only the non-empty grant lists in object_permission", () => {
|
||||
expect(
|
||||
buildOrgCreateBody({
|
||||
...emptyOrgFormValues,
|
||||
organization_alias: "acme",
|
||||
mcp: { servers: [], accessGroups: [], toolsets: ["ts-1"] },
|
||||
}).object_permission,
|
||||
).toStrictEqual({ mcp_toolsets: ["ts-1"] });
|
||||
});
|
||||
|
||||
it("parses metadata into an object instead of sending the raw string", () => {
|
||||
expect(
|
||||
buildOrgCreateBody({ ...emptyOrgFormValues, organization_alias: "acme", metadata: '{"a": 1}' }).metadata,
|
||||
).toStrictEqual({ a: 1 });
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,45 @@
|
|||
import { z } from "zod/v4";
|
||||
|
||||
import type { components } from "@/lib/http/schema";
|
||||
|
||||
import type { OrgSettingsFormValues } from "../org-settings/schema";
|
||||
|
||||
export type OrgCreateBody = components["schemas"]["NewOrganizationRequest"];
|
||||
|
||||
export const emptyOrgFormValues: OrgSettingsFormValues = {
|
||||
organization_alias: "",
|
||||
models: [],
|
||||
max_budget: "",
|
||||
budget_duration: "",
|
||||
tpm_limit: "",
|
||||
rpm_limit: "",
|
||||
vector_stores: [],
|
||||
mcp: { servers: [], accessGroups: [], toolsets: [] },
|
||||
metadata: "",
|
||||
};
|
||||
|
||||
const metadataRecordSchema = z.record(z.string(), z.unknown());
|
||||
|
||||
const objectPermissionFromValues = (values: OrgSettingsFormValues): OrgCreateBody["object_permission"] => {
|
||||
const grants = {
|
||||
...(values.vector_stores.length > 0 && { vector_stores: values.vector_stores }),
|
||||
...(values.mcp.servers.length > 0 && { mcp_servers: values.mcp.servers }),
|
||||
...(values.mcp.accessGroups.length > 0 && { mcp_access_groups: values.mcp.accessGroups }),
|
||||
...(values.mcp.toolsets.length > 0 && { mcp_toolsets: values.mcp.toolsets }),
|
||||
};
|
||||
return Object.keys(grants).length > 0 ? grants : undefined;
|
||||
};
|
||||
|
||||
export const buildOrgCreateBody = (values: OrgSettingsFormValues): OrgCreateBody => {
|
||||
const objectPermission = objectPermissionFromValues(values);
|
||||
return {
|
||||
organization_alias: values.organization_alias,
|
||||
models: values.models,
|
||||
...(values.max_budget.trim() !== "" && { max_budget: Number(values.max_budget) }),
|
||||
...(values.tpm_limit.trim() !== "" && { tpm_limit: Number(values.tpm_limit) }),
|
||||
...(values.rpm_limit.trim() !== "" && { rpm_limit: Number(values.rpm_limit) }),
|
||||
...(values.budget_duration !== "" && { budget_duration: values.budget_duration }),
|
||||
...(values.metadata.trim() !== "" && { metadata: metadataRecordSchema.parse(JSON.parse(values.metadata)) }),
|
||||
...(objectPermission !== undefined && { object_permission: objectPermission }),
|
||||
};
|
||||
};
|
||||
|
|
@ -22,9 +22,9 @@ import { fetchClient } from "@/lib/http/api";
|
|||
import { buildOrgPatch, orgToForm, type OrgPatchBody } from "./mapper";
|
||||
import { orgSettingsSchema } from "./schema";
|
||||
|
||||
const NO_RESET = "never";
|
||||
export const NO_RESET = "never";
|
||||
|
||||
const BUDGET_DURATION_OPTIONS = [
|
||||
export const BUDGET_DURATION_OPTIONS = [
|
||||
{ value: NO_RESET, label: "No reset" },
|
||||
{ value: "24h", label: "daily" },
|
||||
{ value: "7d", label: "weekly" },
|
||||
|
|
@ -102,6 +102,7 @@ export const OrgSettingsForm = ({
|
|||
<FormField control={form.control} name="budget_duration" label="Reset Budget">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
<Select
|
||||
items={BUDGET_DURATION_OPTIONS}
|
||||
value={value === "" ? NO_RESET : value}
|
||||
onValueChange={(selected) => onChange(selected === NO_RESET ? "" : selected)}
|
||||
>
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import {
|
|||
import { useGuardrails, GuardrailListItem } from "@/app/(dashboard)/hooks/guardrails/useGuardrails";
|
||||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
import { mapEmptyStringToNull } from "@/utils/keyUpdateUtils";
|
||||
import type { ObjectPermission } from "@/components/object_permission_types";
|
||||
import { isProxyAdminRole } from "@/utils/roles";
|
||||
import {
|
||||
EditOutlined,
|
||||
|
|
@ -118,17 +119,7 @@ export interface TeamData {
|
|||
router_settings?: Record<string, any>;
|
||||
guardrails?: string[];
|
||||
policies?: string[];
|
||||
object_permission?: {
|
||||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
mcp_toolsets?: string[];
|
||||
vector_stores: string[];
|
||||
agents?: string[];
|
||||
agent_access_groups?: string[];
|
||||
search_tools?: string[];
|
||||
};
|
||||
object_permission?: ObjectPermission | null;
|
||||
team_member_budget_table: {
|
||||
max_budget: number;
|
||||
budget_duration: string;
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue