mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
Merge branch 'litellm_internal_staging' into litellm_tool_spend_click_to_logs
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
36dbb0ecf6
104 changed files with 5040 additions and 2472 deletions
|
|
@ -0,0 +1,2 @@
|
|||
-- CreateIndex
|
||||
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogToolIndex_start_time_idx" ON "LiteLLM_SpendLogToolIndex"("start_time");
|
||||
|
|
@ -1094,6 +1094,7 @@ model LiteLLM_SpendLogToolIndex {
|
|||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
@@index([start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1455,6 +1455,7 @@ SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEA
|
|||
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
|
||||
os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5)
|
||||
)
|
||||
TOOL_SPEND_MAX_WINDOW_DAYS = 30
|
||||
SPEND_LOG_PARTITION_INTERVAL = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day")
|
||||
SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7))
|
||||
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -17,6 +17,10 @@ from litellm.llms.bedrock.common_utils import build_mantle_messages_url
|
|||
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
|
||||
AmazonAnthropicClaudeMessagesConfig,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
AnthropicUsage,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -103,6 +107,25 @@ class AmazonMantleMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
|
|||
)
|
||||
return {**request, "model": model_id, **stream_fields}
|
||||
|
||||
def transform_anthropic_messages_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> AnthropicMessagesResponse:
|
||||
response = super().transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
existing_usage: AnthropicUsage = response.get("usage") or AnthropicUsage()
|
||||
normalized_usage: AnthropicUsage = {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
**existing_usage,
|
||||
}
|
||||
return {**response, "usage": normalized_usage}
|
||||
|
||||
def get_async_streaming_response_iterator(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -27417,7 +27417,7 @@
|
|||
},
|
||||
"/v1/tool/spend": {
|
||||
"get": {
|
||||
"description": "Spend attributed to each tool over a date range, for the Cost Optimization dashboard.\n\nJoins ``LiteLLM_SpendLogToolIndex`` (which tool names ran on which request) to\n``LiteLLM_SpendLogs`` (what the request cost). A request that used multiple tools\ncounts its full spend toward each of those tools, so per-tool numbers are\nattributions. ``total_spend`` is the deduplicated spend of every request that\ncalled at least one tool in the window, so it never double counts.",
|
||||
"description": "Spend attributed to each tool over a date range, for the Cost Optimization dashboard.\n\nJoins ``LiteLLM_SpendLogToolIndex`` (which tool names ran on which request) to\n``LiteLLM_SpendLogs`` (what the request cost). A request that used multiple tools\ncounts its full spend toward each of those tools, so per-tool numbers are\nattributions. ``total_spend`` is the deduplicated spend of every request that\ncalled at least one tool in the window, so it never double counts.\n\n``start_date`` is clamped to at most 30 days before ``end_date`` (serving up to\n31 calendar dates inclusive, the same width as the endpoint's default window):\na wider requested range is clamped, and the response's ``start_date`` reflects\nthe effective window actually served.",
|
||||
"operationId": "get_tool_spend_v1_tool_spend_get",
|
||||
"parameters": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -2320,6 +2320,9 @@ async def _run_centralized_common_checks(
|
|||
None if isinstance(global_spend_result, BaseException) else global_spend_result
|
||||
)
|
||||
|
||||
if user_api_key_auth_obj.org_id is None and team_object is not None and team_object.organization_id is not None:
|
||||
user_api_key_auth_obj.org_id = team_object.organization_id
|
||||
|
||||
# common_checks identifies admin via user_object, not the token
|
||||
# (non_proxy_admin_allowed_routes_check). JWT admin shortcut and
|
||||
# master_key tokens get admin from the token; the DB row for the
|
||||
|
|
|
|||
|
|
@ -38,11 +38,37 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
update_batch_in_database,
|
||||
)
|
||||
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
|
||||
from litellm.repositories.table_repositories import ManagedFileRepository
|
||||
from litellm.types.llms.openai import LiteLLMBatchCreateRequest
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str | None":
|
||||
"""Resolve a managed (unified) input_file_id to its backend storage_url.
|
||||
|
||||
Provider batch handlers (e.g. Vertex AI, which parses a `publishers/`
|
||||
segment out of the file URI) need a real storage location; the opaque
|
||||
unified token crashes them. Returns None whenever a storage_url cannot be
|
||||
produced (no database, lookup error, no managed-file row, or a row without
|
||||
a storage_url yet) so callers fall back to dispatching the original id,
|
||||
which the managed-files deployment hook still maps. This adds resolution
|
||||
without changing behavior on any path that did not resolve before.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
return None
|
||||
try:
|
||||
db_file = await ManagedFileRepository(prisma_client).table.find_first(where={"unified_file_id": input_file_id})
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("create_batch: managed file lookup failed for %s: %s", input_file_id, e)
|
||||
return None
|
||||
if db_file is None:
|
||||
return None
|
||||
return db_file.storage_url or None
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{provider}/v1/batches",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -224,6 +250,11 @@ async def create_batch(
|
|||
)
|
||||
model = target_model_names[0]
|
||||
_create_batch_data["model"] = model
|
||||
|
||||
resolved_storage_url = await _resolve_managed_input_file_storage_url(input_file_id)
|
||||
if resolved_storage_url is not None:
|
||||
_create_batch_data["input_file_id"] = resolved_storage_url
|
||||
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -73,30 +73,41 @@ class SpendLogCleanup:
|
|||
)
|
||||
return False
|
||||
|
||||
async def _delete_old_logs(self, prisma_client: PrismaClient, cutoff_date: datetime) -> int:
|
||||
async def _delete_old_rows_batched(
|
||||
self,
|
||||
prisma_client: PrismaClient,
|
||||
cutoff_date: datetime,
|
||||
table_name: str,
|
||||
key_columns: tuple[str, ...],
|
||||
time_column: str,
|
||||
) -> int:
|
||||
"""
|
||||
Helper method to delete old logs in batches.
|
||||
Returns the total number of logs deleted.
|
||||
Helper method to delete a table's rows older than the cutoff in batches.
|
||||
Returns the total number of rows deleted.
|
||||
"""
|
||||
key_list = ", ".join(f'"{col}"' for col in key_columns)
|
||||
delete_sql = f"""
|
||||
DELETE FROM "{table_name}"
|
||||
WHERE ({key_list}) IN (
|
||||
SELECT {key_list} FROM "{table_name}"
|
||||
WHERE "{time_column}" < $1::timestamptz
|
||||
LIMIT $2
|
||||
)
|
||||
"""
|
||||
total_deleted = 0
|
||||
run_count = 0
|
||||
consecutive_failures = 0
|
||||
while True:
|
||||
if run_count > SPEND_LOG_RUN_LOOPS:
|
||||
verbose_proxy_logger.info("Max logs deleted - 1,00,000, rest of the logs will be deleted in next run")
|
||||
verbose_proxy_logger.info(
|
||||
"Max batches reached for %s cleanup, remaining rows will be deleted in next run", table_name
|
||||
)
|
||||
break
|
||||
# Step 1: Find logs and delete them in one go without fetching to application
|
||||
# Step 1: Find rows and delete them in one go without fetching to application
|
||||
# Delete in batches, limited by self.batch_size
|
||||
try:
|
||||
deleted_result = await prisma_client.db.execute_raw(
|
||||
"""
|
||||
DELETE FROM "LiteLLM_SpendLogs"
|
||||
WHERE ("request_id", "startTime") IN (
|
||||
SELECT "request_id", "startTime" FROM "LiteLLM_SpendLogs"
|
||||
WHERE "startTime" < $1::timestamptz
|
||||
LIMIT $2
|
||||
)
|
||||
""",
|
||||
delete_sql,
|
||||
cutoff_date,
|
||||
self.batch_size,
|
||||
)
|
||||
|
|
@ -105,9 +116,10 @@ class SpendLogCleanup:
|
|||
# the whole run — subsequent batches may still succeed.
|
||||
consecutive_failures += 1
|
||||
verbose_proxy_logger.exception(
|
||||
"Spend log cleanup batch failed "
|
||||
"%s cleanup batch failed "
|
||||
"(run_count=%d, consecutive_failures=%d, batch_size=%d, "
|
||||
"cutoff=%s, total_deleted_so_far=%d): %s: %s",
|
||||
table_name,
|
||||
run_count,
|
||||
consecutive_failures,
|
||||
self.batch_size,
|
||||
|
|
@ -118,8 +130,8 @@ class SpendLogCleanup:
|
|||
)
|
||||
if consecutive_failures >= SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES:
|
||||
verbose_proxy_logger.error(
|
||||
"Aborting spend log cleanup after %d consecutive batch "
|
||||
"failures; total deleted before abort: %d",
|
||||
"Aborting %s cleanup after %d consecutive batch failures; total deleted before abort: %d",
|
||||
table_name,
|
||||
consecutive_failures,
|
||||
total_deleted,
|
||||
)
|
||||
|
|
@ -134,15 +146,15 @@ class SpendLogCleanup:
|
|||
deleted_count = deleted_result
|
||||
else:
|
||||
verbose_proxy_logger.error(
|
||||
f"Unexpected execute_raw return type for spend log cleanup: {type(deleted_result)}; "
|
||||
f"Unexpected execute_raw return type for {table_name} cleanup: {type(deleted_result)}; "
|
||||
"aborting cleanup to avoid infinite loop"
|
||||
)
|
||||
break
|
||||
|
||||
verbose_proxy_logger.info(f"Deleted {deleted_count} logs in this batch")
|
||||
verbose_proxy_logger.info(f"Deleted {deleted_count} {table_name} rows in this batch")
|
||||
|
||||
if deleted_count == 0:
|
||||
verbose_proxy_logger.info(f"No more logs to delete. Total deleted: {total_deleted}")
|
||||
verbose_proxy_logger.info(f"No more {table_name} rows to delete. Total deleted: {total_deleted}")
|
||||
break
|
||||
|
||||
total_deleted += deleted_count
|
||||
|
|
@ -153,6 +165,26 @@ class SpendLogCleanup:
|
|||
|
||||
return total_deleted
|
||||
|
||||
async def _delete_old_logs(self, prisma_client: PrismaClient, cutoff_date: datetime) -> int:
|
||||
return await self._delete_old_rows_batched(
|
||||
prisma_client,
|
||||
cutoff_date,
|
||||
table_name="LiteLLM_SpendLogs",
|
||||
key_columns=("request_id", "startTime"),
|
||||
time_column="startTime",
|
||||
)
|
||||
|
||||
async def _delete_old_tool_index_rows(self, prisma_client: PrismaClient, cutoff_date: datetime) -> int:
|
||||
# SpendLogToolIndex rows are derived from spend logs, so they expire on the
|
||||
# same cutoff; rows older than retention point at already-deleted logs.
|
||||
return await self._delete_old_rows_batched(
|
||||
prisma_client,
|
||||
cutoff_date,
|
||||
table_name="LiteLLM_SpendLogToolIndex",
|
||||
key_columns=("request_id", "tool_name"),
|
||||
time_column="start_time",
|
||||
)
|
||||
|
||||
async def cleanup_old_spend_logs(self, prisma_client: PrismaClient) -> None:
|
||||
"""
|
||||
Main cleanup function. Deletes old spend logs in batches.
|
||||
|
|
@ -209,6 +241,9 @@ class SpendLogCleanup:
|
|||
total_deleted = await self._delete_old_logs(prisma_client, cutoff_date)
|
||||
verbose_proxy_logger.info(f"Deleted {total_deleted} logs")
|
||||
|
||||
index_deleted = await self._delete_old_tool_index_rows(prisma_client, cutoff_date)
|
||||
verbose_proxy_logger.info(f"Deleted {index_deleted} expired tool index rows")
|
||||
|
||||
except Exception as e:
|
||||
# .exception() captures the traceback; str(e) alone on a Prisma/DB
|
||||
# timeout is often empty and gives operators no signal to diagnose.
|
||||
|
|
|
|||
|
|
@ -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 [
|
||||
|
|
@ -410,6 +412,19 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
)
|
||||
if key in body
|
||||
}
|
||||
tokens_before = stats.get("tokens_before")
|
||||
tokens_after = stats.get("tokens_after")
|
||||
if (
|
||||
"tokens_saved" not in stats
|
||||
and isinstance(tokens_before, (int, float))
|
||||
and not isinstance(tokens_before, bool)
|
||||
and isinstance(tokens_after, (int, float))
|
||||
and not isinstance(tokens_after, bool)
|
||||
):
|
||||
# Spend tracking (extract_compression_saved_tokens) reads only
|
||||
# tokens_saved, which the live compression service omits; derive it
|
||||
# so savings are counted, but let a service-sent value win.
|
||||
stats["tokens_saved"] = tokens_before - tokens_after
|
||||
return filtered, True, stats
|
||||
|
||||
async def _call_retrieve(self, hash_value: str, query: str | None = None) -> str:
|
||||
|
|
@ -481,7 +496,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 +522,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"),
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import TOOL_SPEND_MAX_WINDOW_DAYS
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
|
||||
|
|
@ -209,6 +210,11 @@ async def get_tool_spend(
|
|||
counts its full spend toward each of those tools, so per-tool numbers are
|
||||
attributions. ``total_spend`` is the deduplicated spend of every request that
|
||||
called at least one tool in the window, so it never double counts.
|
||||
|
||||
``start_date`` is clamped to at most 30 days before ``end_date`` (serving up to
|
||||
31 calendar dates inclusive, the same width as the endpoint's default window):
|
||||
a wider requested range is clamped, and the response's ``start_date`` reflects
|
||||
the effective window actually served.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
|
|
@ -226,9 +232,19 @@ async def get_tool_spend(
|
|||
|
||||
now = datetime.now(timezone.utc)
|
||||
end_day = _parse_day_start(end_date)
|
||||
start_dt = _parse_day_start(start_date) or ((end_day or now) - timedelta(days=30))
|
||||
# Anchor the floor to a midnight so the clamp compares dates with dates:
|
||||
# parsed start_dates are midnight-aligned, and a floor carrying now's
|
||||
# time-of-day would invisibly truncate an explicit start_date to mid-day.
|
||||
today = now.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
window_floor = (end_day or today) - timedelta(days=TOOL_SPEND_MAX_WINDOW_DAYS)
|
||||
start_dt = _parse_day_start(start_date) or window_floor
|
||||
if start_dt < window_floor:
|
||||
start_dt = window_floor
|
||||
end_exclusive = (end_day + timedelta(days=1)) if end_day else now
|
||||
|
||||
# ti.start_time defines the window in both queries; the sl."startTime" bounds
|
||||
# exist only so the planner can use the SpendLogs startTime index, and carry a
|
||||
# 1s margin because the two writers can disagree by ~1ms on the same request.
|
||||
rows = await prisma_client.db.query_raw(
|
||||
"""
|
||||
SELECT to_char(ti.start_time, 'YYYY-MM-DD') AS date,
|
||||
|
|
@ -240,6 +256,8 @@ async def get_tool_spend(
|
|||
JOIN "LiteLLM_SpendLogs" sl ON sl.request_id = ti.request_id
|
||||
WHERE ti.start_time >= ($1::timestamptz AT TIME ZONE 'UTC')
|
||||
AND ti.start_time < ($2::timestamptz AT TIME ZONE 'UTC')
|
||||
AND sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') - interval '1 second'
|
||||
AND sl."startTime" < ($2::timestamptz AT TIME ZONE 'UTC') + interval '1 second'
|
||||
GROUP BY date, ti.tool_name
|
||||
ORDER BY date ASC, spend DESC
|
||||
""",
|
||||
|
|
@ -250,7 +268,9 @@ async def get_tool_spend(
|
|||
"""
|
||||
SELECT COALESCE(SUM(sl.spend), 0)::double precision AS total_spend
|
||||
FROM "LiteLLM_SpendLogs" sl
|
||||
WHERE EXISTS (
|
||||
WHERE sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') - interval '1 second'
|
||||
AND sl."startTime" < ($2::timestamptz AT TIME ZONE 'UTC') + interval '1 second'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM "LiteLLM_SpendLogToolIndex" ti
|
||||
WHERE ti.request_id = sl.request_id
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1094,6 +1094,7 @@ model LiteLLM_SpendLogToolIndex {
|
|||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
@@index([start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1094,6 +1094,7 @@ model LiteLLM_SpendLogToolIndex {
|
|||
|
||||
@@id([request_id, tool_name])
|
||||
@@index([tool_name, start_time])
|
||||
@@index([start_time])
|
||||
}
|
||||
|
||||
// Prompt table for storing prompt configurations
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -2853,3 +2853,77 @@ def test_streaming_function_call_tool_id_for_degenerate_call_id():
|
|||
|
||||
assert stream_tool_id("fc_unique_abc123", "call_0") == "fc_unique_abc123"
|
||||
assert stream_tool_id("fc_2", "call_tokyo") == "call_tokyo"
|
||||
|
||||
|
||||
|
||||
@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.
|
||||
|
|
|
|||
|
|
@ -399,6 +399,103 @@ async def test_mantle_anthropic_messages_sends_workspace_header_and_clean_body()
|
|||
assert "aws_bedrock_project_id" not in requests[0]["body"]
|
||||
|
||||
|
||||
def _usageless_anthropic_response(url: str) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "msg_classifier",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "anthropic.claude-opus-4-8",
|
||||
"content": [{"type": "text", "text": "safe"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
},
|
||||
request=httpx.Request("POST", url),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_anthropic_messages_backfills_missing_usage():
|
||||
"""
|
||||
Regression for LIT-4758: a Mantle non-streaming response with no `usage`
|
||||
object must not reach the client usage-less, or Claude Code's auto-mode
|
||||
classifier crashes on `usage.input_tokens`.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
async def mock_post(self, url, data=None, headers=None, **kwargs):
|
||||
return _usageless_anthropic_response(str(url))
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new=mock_post,
|
||||
):
|
||||
response = await litellm.anthropic_messages(
|
||||
model="bedrock/mantle/anthropic.claude-opus-4-8",
|
||||
messages=[{"role": "user", "content": "is `Bash(ls)` safe?"}],
|
||||
max_tokens=10,
|
||||
aws_access_key_id="fake-key",
|
||||
aws_secret_access_key="fake-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
finally:
|
||||
await litellm.close_litellm_async_clients()
|
||||
|
||||
assert response["usage"]["input_tokens"] == 0
|
||||
assert response["usage"]["output_tokens"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_anthropic_messages_preserves_upstream_usage():
|
||||
"""Backfill must not clobber a usage object the upstream did return."""
|
||||
import litellm
|
||||
|
||||
def _response_with_usage(url: str) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "msg_test",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "anthropic.claude-opus-4-8",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {
|
||||
"input_tokens": 42,
|
||||
"output_tokens": 7,
|
||||
"cache_read_input_tokens": 5,
|
||||
},
|
||||
},
|
||||
request=httpx.Request("POST", url),
|
||||
)
|
||||
|
||||
async def mock_post(self, url, data=None, headers=None, **kwargs):
|
||||
return _response_with_usage(str(url))
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new=mock_post,
|
||||
):
|
||||
response = await litellm.anthropic_messages(
|
||||
model="bedrock/mantle/anthropic.claude-opus-4-8",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
max_tokens=10,
|
||||
aws_access_key_id="fake-key",
|
||||
aws_secret_access_key="fake-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
finally:
|
||||
await litellm.close_litellm_async_clients()
|
||||
|
||||
assert response["usage"]["input_tokens"] == 42
|
||||
assert response["usage"]["output_tokens"] == 7
|
||||
assert response["usage"]["cache_read_input_tokens"] == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mantle_anthropic_messages_routes_to_vpc_api_base():
|
||||
import litellm
|
||||
|
|
|
|||
|
|
@ -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({}) == ""
|
||||
|
||||
|
|
|
|||
|
|
@ -3796,6 +3796,167 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on
|
|||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"key_org_id,team_org_id,expected_org_id",
|
||||
[
|
||||
(None, "org-from-team", "org-from-team"),
|
||||
("org-pinned-on-key", "org-from-team", "org-pinned-on-key"),
|
||||
(None, None, None),
|
||||
],
|
||||
)
|
||||
async def test_centralized_common_checks_backfills_org_id_from_team(key_org_id, team_org_id, expected_org_id):
|
||||
"""LIT-4688 regression: a key minted without an organization_id but attached
|
||||
to an org-linked team must leave auth with org_id set from the team, so the
|
||||
spend writer (which reads user_api_key_dict.org_id, no team fallback)
|
||||
credits the org and the org budget cap can actually trip. A key with an
|
||||
explicitly pinned org_id must win over the team's org."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
|
||||
token = UserAPIKeyAuth(api_key="sk-test", user_id="u", team_id="t1", org_id=key_org_id)
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
fetched_team = LiteLLM_TeamTableCachedObj(team_id="t1", organization_id=team_org_id)
|
||||
|
||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
org_id_seen_by_common_checks = []
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=fetched_team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=lambda **kw: org_id_seen_by_common_checks.append(kw["valid_token"].org_id),
|
||||
) as mock_checks,
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
mock_checks.assert_awaited_once()
|
||||
assert token.org_id == expected_org_id
|
||||
assert org_id_seen_by_common_checks == [expected_org_id]
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_session_token_org_backfilled_from_team(monkeypatch):
|
||||
"""LIT-4688 root cause: CLI session tokens (from /sso/cli/poll) are minted
|
||||
with a real team_id but no org_id, and their auth path decrypts the blob
|
||||
without the combined_view team join, so their spend never reached the org.
|
||||
The centralized-checks backfill must complete the credential from the team
|
||||
the same way the SQL view does for DB keys."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-lit4688")
|
||||
|
||||
cli_user = LiteLLM_UserTable(user_id="cli-user", user_role="internal_user", teams=["t-cli"], models=[])
|
||||
blob = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id="t-cli", team_alias="cli-team")
|
||||
token = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(blob)
|
||||
assert token is not None
|
||||
assert token.is_session_token is True
|
||||
assert token.team_id == "t-cli"
|
||||
assert token.org_id is None
|
||||
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
org_linked_team = LiteLLM_TeamTableCachedObj(team_id="t-cli", organization_id="org-infoops")
|
||||
|
||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=org_linked_team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert token.org_id == "org-infoops"
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_centralized_common_checks_org_backfill_survives_team_fetch_failure():
|
||||
"""When the team DB fetch fails, the token-derived fallback team carries no
|
||||
organization_id, so the backfill must leave org_id as None rather than
|
||||
crash or mis-attribute."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token = UserAPIKeyAuth(api_key="sk-test", user_id="u", team_id="t1")
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=Exception("DB down"),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_checks,
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
mock_checks.assert_awaited_once()
|
||||
assert token.org_id is None
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_master_key_auth_substitutes_alias_for_api_key():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -76,9 +76,7 @@ CREDS: Dict[str, Dict[str, str]] = {
|
|||
}
|
||||
|
||||
# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123".
|
||||
AZURE_FILE_ID = encode_file_id_with_model(
|
||||
"file-original123", "azure/gpt-4o", id_type="file"
|
||||
)
|
||||
AZURE_FILE_ID = encode_file_id_with_model("file-original123", "azure/gpt-4o", id_type="file")
|
||||
|
||||
|
||||
def make_batch(
|
||||
|
|
@ -166,9 +164,7 @@ def harness():
|
|||
|
||||
router = MagicMock(spec=Router)
|
||||
router.acreate_batch = AsyncMock(return_value=make_batch())
|
||||
router.get_deployment_credentials_with_provider = MagicMock(
|
||||
side_effect=_creds_lookup
|
||||
)
|
||||
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
|
||||
|
||||
read_body = AsyncMock(side_effect=lambda request: body_holder["body"])
|
||||
pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock()))
|
||||
|
|
@ -186,11 +182,7 @@ def harness():
|
|||
pre_call,
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers
|
||||
)
|
||||
)
|
||||
stack.enter_context(patch.object(ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers))
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
endpoints,
|
||||
|
|
@ -200,14 +192,13 @@ def harness():
|
|||
)
|
||||
stack.enter_context(patch.object(endpoints, "is_known_model", is_known_model))
|
||||
stack.enter_context(patch.object(litellm, "acreate_batch", litellm_acreate))
|
||||
stack.enter_context(
|
||||
patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False)
|
||||
)
|
||||
stack.enter_context(patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False))
|
||||
stack.enter_context(patch.object(proxy_server, "llm_router", router))
|
||||
stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging))
|
||||
stack.enter_context(patch.object(proxy_server, "general_settings", {}))
|
||||
stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock()))
|
||||
stack.enter_context(patch.object(proxy_server, "version", "test-version"))
|
||||
stack.enter_context(patch.object(proxy_server, "prisma_client", None))
|
||||
|
||||
h = Harness(
|
||||
body=body_holder,
|
||||
|
|
@ -283,9 +274,7 @@ async def test_create__model_encoded_file_id(harness):
|
|||
}
|
||||
|
||||
# 4. OUTPUT SHAPE - ids re-encoded with the model; input_file_id restored.
|
||||
assert resp.id == encode_file_id_with_model(
|
||||
"batch-provider-id", "azure/gpt-4o", id_type="batch"
|
||||
)
|
||||
assert resp.id == encode_file_id_with_model("batch-provider-id", "azure/gpt-4o", id_type="batch")
|
||||
assert resp.input_file_id == AZURE_FILE_ID
|
||||
|
||||
|
||||
|
|
@ -307,12 +296,8 @@ async def test_create__model_encoded_file_id__encodes_output_and_error_ids(harne
|
|||
|
||||
resp = await call_create(harness)
|
||||
|
||||
assert resp.output_file_id == encode_file_id_with_model(
|
||||
"file-out-raw", "azure/gpt-4o"
|
||||
)
|
||||
assert resp.error_file_id == encode_file_id_with_model(
|
||||
"file-err-raw", "azure/gpt-4o"
|
||||
)
|
||||
assert resp.output_file_id == encode_file_id_with_model("file-out-raw", "azure/gpt-4o")
|
||||
assert resp.error_file_id == encode_file_id_with_model("file-err-raw", "azure/gpt-4o")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -358,9 +343,7 @@ async def test_create__model_from_body(harness):
|
|||
payload = harness.acreate_kwargs()
|
||||
assert payload["custom_llm_provider"] == "vertex_ai"
|
||||
assert payload["input_file_id"] == "file-plain"
|
||||
assert resp.id == encode_file_id_with_model(
|
||||
"batch-provider-id", "vertex-model", id_type="batch"
|
||||
)
|
||||
assert resp.id == encode_file_id_with_model("batch-provider-id", "vertex-model", id_type="batch")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -495,10 +478,9 @@ async def test_create__unified_file_id_single_model(harness):
|
|||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"
|
||||
), patch.object(
|
||||
endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]),
|
||||
):
|
||||
resp = await call_create(harness)
|
||||
|
||||
|
|
@ -522,10 +504,9 @@ async def test_create__unified_file_id_not_exactly_one_model_400(harness, models
|
|||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"
|
||||
), patch.object(
|
||||
endpoints, "get_models_from_unified_file_id", return_value=models
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=models),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_create(harness)
|
||||
|
|
@ -535,6 +516,177 @@ async def test_create__unified_file_id_not_exactly_one_model_400(harness, models
|
|||
harness.litellm_acreate.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_file_id_resolves_real_storage_url(harness):
|
||||
"""A base64 unified_file_id is a LiteLLM-internal token, not a real
|
||||
provider-side file reference (e.g. Vertex AI's batch transformation parses
|
||||
a `publishers/` segment out of the file URI and crashes on the opaque
|
||||
base64 string). The real backend location (`storage_url`) must be looked
|
||||
up from LiteLLM_ManagedFileTable and substituted before dispatch.
|
||||
|
||||
Regression lock on the lookup key: LiteLLM_ManagedFileTable.unified_file_id
|
||||
stores the raw base64 file id (see schema.prisma and the enterprise
|
||||
managed-files hook, which queries with the raw id), NOT the decoded
|
||||
litellm_proxy:... string. Querying with the decoded string never matches
|
||||
and silently falls back."""
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": "litellm_proxy_unified_id",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
|
||||
fake_db_file = MagicMock(
|
||||
storage_url="gs://bucket/litellm-vertex-files/publishers/google/models/gemini-2.0/abc",
|
||||
)
|
||||
find_first = AsyncMock(return_value=fake_db_file)
|
||||
fake_repo_instance = MagicMock()
|
||||
fake_repo_instance.table.find_first = find_first
|
||||
fake_repo_cls = MagicMock(return_value=fake_repo_instance)
|
||||
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]),
|
||||
patch.object(proxy_server, "prisma_client", MagicMock()),
|
||||
patch.object(endpoints, "ManagedFileRepository", fake_repo_cls),
|
||||
):
|
||||
resp = await call_create(harness)
|
||||
|
||||
assert harness.router_kwargs()["input_file_id"] == fake_db_file.storage_url
|
||||
find_first.assert_awaited_once_with(where={"unified_file_id": "litellm_proxy_unified_id"})
|
||||
assert resp.input_file_id == "litellm_proxy_unified_id"
|
||||
assert resp._hidden_params["unified_file_id"] == "unified-xyz"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_file_id_db_error_falls_back_to_raw_id(harness):
|
||||
"""Resolution is additive and best-effort: a lookup error leaves the id
|
||||
unresolved and dispatch falls back to the original id, exactly as before
|
||||
this change (the managed-files deployment hook still maps it). No new
|
||||
failure mode is introduced."""
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": "litellm_proxy_unified_id",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
|
||||
find_first = AsyncMock(side_effect=Exception("db unavailable"))
|
||||
fake_repo_instance = MagicMock()
|
||||
fake_repo_instance.table.find_first = find_first
|
||||
fake_repo_cls = MagicMock(return_value=fake_repo_instance)
|
||||
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]),
|
||||
patch.object(proxy_server, "prisma_client", MagicMock()),
|
||||
patch.object(endpoints, "ManagedFileRepository", fake_repo_cls),
|
||||
):
|
||||
await call_create(harness)
|
||||
|
||||
assert harness.router_kwargs()["input_file_id"] == "litellm_proxy_unified_id"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__multi_model_unified_file_with_loadbalancing_keeps_router_branch(harness):
|
||||
"""Regression guard: a multi-model managed file dispatched with an explicit
|
||||
router model under load balancing must keep taking the load-balanced router
|
||||
branch, exactly as on the base revision, where the managed-files deployment
|
||||
hook remaps the unified id per model. Routing it into the unified branch
|
||||
instead would trip that branch's "exactly one model" 400 and break a path
|
||||
that works today, so the unified-file resolution must not steal the
|
||||
load-balanced branch."""
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": "litellm_proxy_unified_id",
|
||||
"model": "vertex-model",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
harness.is_known_model.return_value = True
|
||||
|
||||
with (
|
||||
patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True),
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["model-a", "model-b"]),
|
||||
):
|
||||
await call_create(harness)
|
||||
|
||||
assert harness.router_acreate.call_count == 1
|
||||
assert harness.router_kwargs()["input_file_id"] == "litellm_proxy_unified_id"
|
||||
harness.litellm_acreate.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_file_id_missing_row_falls_back_to_raw_id(harness):
|
||||
"""Resolution is additive: when no managed-file row exists there is nothing
|
||||
to substitute, so dispatch falls back to the original id exactly as before
|
||||
this change (the managed-files deployment hook still maps it). No new
|
||||
failure mode is introduced for this case."""
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": "litellm_proxy_unified_id",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
|
||||
find_first = AsyncMock(return_value=None)
|
||||
fake_repo_instance = MagicMock()
|
||||
fake_repo_instance.table.find_first = find_first
|
||||
fake_repo_cls = MagicMock(return_value=fake_repo_instance)
|
||||
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]),
|
||||
patch.object(proxy_server, "prisma_client", MagicMock()),
|
||||
patch.object(endpoints, "ManagedFileRepository", fake_repo_cls),
|
||||
):
|
||||
await call_create(harness)
|
||||
|
||||
assert harness.router_kwargs()["input_file_id"] == "litellm_proxy_unified_id"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__unified_file_id_legacy_row_without_storage_url_dispatches_raw(
|
||||
harness,
|
||||
):
|
||||
"""A managed file whose row predates the storage_url column still dispatches
|
||||
the original id (the managed-files deployment hook maps it); the row exists,
|
||||
so this is not the missing-row fail-closed case."""
|
||||
set_body(
|
||||
harness,
|
||||
{
|
||||
"input_file_id": "litellm_proxy_unified_id",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
|
||||
fake_db_file = MagicMock(storage_url=None)
|
||||
find_first = AsyncMock(return_value=fake_db_file)
|
||||
fake_repo_instance = MagicMock()
|
||||
fake_repo_instance.table.find_first = find_first
|
||||
fake_repo_cls = MagicMock(return_value=fake_repo_instance)
|
||||
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gemini-2.0"]),
|
||||
patch.object(proxy_server, "prisma_client", MagicMock()),
|
||||
patch.object(endpoints, "ManagedFileRepository", fake_repo_cls),
|
||||
):
|
||||
await call_create(harness)
|
||||
|
||||
assert harness.router_kwargs()["input_file_id"] == "litellm_proxy_unified_id"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create__model_encoded_beats_unified(harness):
|
||||
"""Precedence row: a file id that is BOTH model-encoded and (pretend) unified
|
||||
|
|
@ -547,10 +699,9 @@ async def test_create__model_encoded_beats_unified(harness):
|
|||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"
|
||||
), patch.object(
|
||||
endpoints, "get_models_from_unified_file_id", return_value=["something-else"]
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["something-else"]),
|
||||
):
|
||||
await call_create(harness)
|
||||
|
||||
|
|
@ -579,9 +730,7 @@ async def test_create__loadbalancing_routes_to_router(harness):
|
|||
with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True):
|
||||
await call_create(harness)
|
||||
|
||||
harness.is_known_model.assert_called_once_with(
|
||||
model="lb-model", llm_router=harness.router
|
||||
)
|
||||
harness.is_known_model.assert_called_once_with(model="lb-model", llm_router=harness.router)
|
||||
assert harness.router_acreate.call_count == 1
|
||||
harness.litellm_acreate.assert_not_called()
|
||||
harness.creds_resolver.assert_not_called()
|
||||
|
|
@ -630,9 +779,7 @@ async def test_create__team_expiry_injected(harness):
|
|||
},
|
||||
)
|
||||
|
||||
await call_create(
|
||||
harness, user=_user_with_expiry({"anchor": "created_at", "seconds": 3600})
|
||||
)
|
||||
await call_create(harness, user=_user_with_expiry({"anchor": "created_at", "seconds": 3600}))
|
||||
|
||||
assert harness.acreate_kwargs()["output_expires_after"] == {
|
||||
"anchor": "created_at",
|
||||
|
|
@ -738,12 +885,7 @@ async def test_create__exception_calls_failure_hook(harness):
|
|||
await call_create(harness)
|
||||
|
||||
harness.logging.post_call_failure_hook.assert_called_once()
|
||||
assert (
|
||||
harness.logging.post_call_failure_hook.call_args.kwargs[
|
||||
"original_exception"
|
||||
].args[0]
|
||||
== "provider boom"
|
||||
)
|
||||
assert harness.logging.post_call_failure_hook.call_args.kwargs["original_exception"].args[0] == "provider boom"
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
|
|
@ -770,9 +912,7 @@ async def test_create__exception_calls_failure_hook(harness):
|
|||
# A real model-encoded BATCH id: decodes to "azure/gpt-4o", strips to
|
||||
# "batch_orig123". Distinct from AZURE_FILE_ID so retrieve tests can't pass by
|
||||
# accidentally reusing the create fixture's value.
|
||||
AZURE_BATCH_ID = encode_file_id_with_model(
|
||||
"batch_orig123", "azure/gpt-4o", id_type="batch"
|
||||
)
|
||||
AZURE_BATCH_ID = encode_file_id_with_model("batch_orig123", "azure/gpt-4o", id_type="batch")
|
||||
|
||||
# A realistic decoded unified batch id (what _is_base64_encoded_unified_file_id
|
||||
# returns). model_id / llm_batch_id are parsed out of this by the real helpers.
|
||||
|
|
@ -823,9 +963,7 @@ def retrieve_harness():
|
|||
|
||||
router = MagicMock(spec=Router)
|
||||
router.aretrieve_batch = AsyncMock(return_value=make_batch())
|
||||
router.get_deployment_credentials_with_provider = MagicMock(
|
||||
side_effect=_creds_lookup
|
||||
)
|
||||
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
|
||||
|
||||
pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock()))
|
||||
get_headers = MagicMock(return_value={})
|
||||
|
|
@ -846,11 +984,7 @@ def retrieve_harness():
|
|||
pre_call,
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers
|
||||
)
|
||||
)
|
||||
stack.enter_context(patch.object(ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers))
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
endpoints,
|
||||
|
|
@ -865,24 +999,12 @@ def retrieve_harness():
|
|||
provider_from_query,
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(endpoints, "get_batch_from_database", get_batch_from_db)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(endpoints, "update_batch_in_database", update_batch_in_db)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(endpoints, "resolve_input_file_id_to_unified", resolve_input)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
endpoints, "resolve_output_file_ids_to_unified", resolve_output
|
||||
)
|
||||
)
|
||||
stack.enter_context(patch.object(endpoints, "get_batch_from_database", get_batch_from_db))
|
||||
stack.enter_context(patch.object(endpoints, "update_batch_in_database", update_batch_in_db))
|
||||
stack.enter_context(patch.object(endpoints, "resolve_input_file_id_to_unified", resolve_input))
|
||||
stack.enter_context(patch.object(endpoints, "resolve_output_file_ids_to_unified", resolve_output))
|
||||
stack.enter_context(patch.object(litellm, "aretrieve_batch", litellm_aretrieve))
|
||||
stack.enter_context(
|
||||
patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False)
|
||||
)
|
||||
stack.enter_context(patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False))
|
||||
stack.enter_context(patch.object(proxy_server, "llm_router", router))
|
||||
stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging))
|
||||
stack.enter_context(patch.object(proxy_server, "general_settings", {}))
|
||||
|
|
@ -956,9 +1078,7 @@ async def test_retrieve__model_encoded_id(retrieve_harness):
|
|||
}
|
||||
|
||||
# 4. OUTPUT SHAPE - ids re-encoded with the model for the round-trip.
|
||||
assert resp.id == encode_file_id_with_model(
|
||||
"batch-provider-id", "azure/gpt-4o", id_type="batch"
|
||||
)
|
||||
assert resp.id == encode_file_id_with_model("batch-provider-id", "azure/gpt-4o", id_type="batch")
|
||||
|
||||
# write-back to the managed-object table happened, tagged as a retrieve.
|
||||
assert retrieve_harness.update_batch_in_db.call_count == 1
|
||||
|
|
@ -989,12 +1109,8 @@ async def test_retrieve__model_encoded_id__encodes_output_and_error_ids(
|
|||
|
||||
resp = await call_retrieve(retrieve_harness, AZURE_BATCH_ID)
|
||||
|
||||
assert resp.output_file_id == encode_file_id_with_model(
|
||||
"file-out-raw", "azure/gpt-4o"
|
||||
)
|
||||
assert resp.error_file_id == encode_file_id_with_model(
|
||||
"file-err-raw", "azure/gpt-4o"
|
||||
)
|
||||
assert resp.output_file_id == encode_file_id_with_model("file-out-raw", "azure/gpt-4o")
|
||||
assert resp.error_file_id == encode_file_id_with_model("file-err-raw", "azure/gpt-4o")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1018,9 +1134,7 @@ async def test_retrieve__model_encoded_beats_loadbalancing(retrieve_harness):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness):
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID
|
||||
):
|
||||
with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID):
|
||||
resp = await call_retrieve(retrieve_harness, "batch-unified-blob")
|
||||
|
||||
# DISPATCH - router fired, direct litellm did not.
|
||||
|
|
@ -1125,9 +1239,7 @@ async def test_retrieve__fallback_provider_precedence_path_over_header(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"status", ["completed", "complete", "failed", "cancelled", "expired"]
|
||||
)
|
||||
@pytest.mark.parametrize("status", ["completed", "complete", "failed", "cancelled", "expired"])
|
||||
async def test_retrieve__db_terminal_state_short_circuits(retrieve_harness, status):
|
||||
# "complete" is the DB-normalized alias of "completed"; it is not a valid
|
||||
# constructor literal but reaches the endpoint via a stored row, so set it
|
||||
|
|
@ -1151,9 +1263,7 @@ async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness)
|
|||
db_response = make_batch(id="batch-from-db", status="completed")
|
||||
retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response)
|
||||
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID
|
||||
):
|
||||
with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID):
|
||||
await call_retrieve(retrieve_harness, "batch-unified-blob")
|
||||
|
||||
# Terminal short-circuit still resolves raw provider file ids to unified.
|
||||
|
|
@ -1186,9 +1296,7 @@ async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harn
|
|||
async def test_retrieve__uses_aretrieve_batch_route_type(retrieve_harness):
|
||||
await call_retrieve(retrieve_harness, "batch-raw-xyz")
|
||||
|
||||
assert (
|
||||
retrieve_harness.pre_call.call_args.kwargs["route_type"] == "aretrieve_batch"
|
||||
)
|
||||
assert retrieve_harness.pre_call.call_args.kwargs["route_type"] == "aretrieve_batch"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1200,9 +1308,7 @@ async def test_retrieve__exception_calls_failure_hook(retrieve_harness):
|
|||
|
||||
retrieve_harness.logging.post_call_failure_hook.assert_called_once()
|
||||
assert (
|
||||
retrieve_harness.logging.post_call_failure_hook.call_args.kwargs[
|
||||
"original_exception"
|
||||
].args[0]
|
||||
retrieve_harness.logging.post_call_failure_hook.call_args.kwargs["original_exception"].args[0]
|
||||
== "provider boom"
|
||||
)
|
||||
|
||||
|
|
@ -1275,9 +1381,7 @@ def list_harness():
|
|||
|
||||
router = MagicMock(spec=Router)
|
||||
router.alist_batches = AsyncMock(return_value=FakeListPage([]))
|
||||
router.get_deployment_credentials_with_provider = MagicMock(
|
||||
side_effect=_creds_lookup
|
||||
)
|
||||
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
|
||||
|
||||
read_body = AsyncMock(side_effect=lambda request: body_holder["body"])
|
||||
pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock()))
|
||||
|
|
@ -1295,11 +1399,7 @@ def list_harness():
|
|||
pre_call,
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers
|
||||
)
|
||||
)
|
||||
stack.enter_context(patch.object(ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers))
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
endpoints,
|
||||
|
|
@ -1432,21 +1532,15 @@ async def test_list__managed_files_beats_model_param(list_harness):
|
|||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_list__model_from_body_routes_and_encodes(list_harness):
|
||||
list_harness.litellm_alist.return_value = FakeListPage(
|
||||
[make_batch(id="batch-1"), make_batch(id="batch-2")]
|
||||
)
|
||||
list_harness.litellm_alist.return_value = FakeListPage([make_batch(id="batch-1"), make_batch(id="batch-2")])
|
||||
|
||||
resp = await call_list(list_harness, body={"model": "azure/gpt-4o"})
|
||||
|
||||
assert list_harness.litellm_alist.call_count == 1
|
||||
list_harness.router_alist.assert_not_called()
|
||||
list_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o")
|
||||
assert resp.data[0].id == encode_file_id_with_model(
|
||||
"batch-1", "azure/gpt-4o", id_type="batch"
|
||||
)
|
||||
assert resp.data[1].id == encode_file_id_with_model(
|
||||
"batch-2", "azure/gpt-4o", id_type="batch"
|
||||
)
|
||||
assert resp.data[0].id == encode_file_id_with_model("batch-1", "azure/gpt-4o", id_type="batch")
|
||||
assert resp.data[1].id == encode_file_id_with_model("batch-2", "azure/gpt-4o", id_type="batch")
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
@ -1577,12 +1671,7 @@ async def test_list__exception_calls_failure_hook(list_harness):
|
|||
await call_list(list_harness)
|
||||
|
||||
list_harness.logging.post_call_failure_hook.assert_called_once()
|
||||
assert (
|
||||
list_harness.logging.post_call_failure_hook.call_args.kwargs[
|
||||
"original_exception"
|
||||
].args[0]
|
||||
== "provider boom"
|
||||
)
|
||||
assert list_harness.logging.post_call_failure_hook.call_args.kwargs["original_exception"].args[0] == "provider boom"
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
|
|
@ -1645,9 +1734,7 @@ def cancel_harness():
|
|||
|
||||
router = MagicMock(spec=Router)
|
||||
router.acancel_batch = AsyncMock(return_value=make_batch())
|
||||
router.get_deployment_credentials_with_provider = MagicMock(
|
||||
side_effect=_creds_lookup
|
||||
)
|
||||
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
|
||||
|
||||
pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock()))
|
||||
# add_litellm_data_to_request is a passthrough that returns the data it got.
|
||||
|
|
@ -1666,11 +1753,7 @@ def cancel_harness():
|
|||
pre_call,
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers
|
||||
)
|
||||
)
|
||||
stack.enter_context(patch.object(ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers))
|
||||
stack.enter_context(
|
||||
patch.object(
|
||||
endpoints,
|
||||
|
|
@ -1685,22 +1768,16 @@ def cancel_harness():
|
|||
provider_from_query,
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(endpoints, "update_batch_in_database", update_batch_in_db)
|
||||
)
|
||||
stack.enter_context(patch.object(endpoints, "update_batch_in_database", update_batch_in_db))
|
||||
stack.enter_context(patch.object(litellm, "acancel_batch", litellm_acancel))
|
||||
stack.enter_context(
|
||||
patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False)
|
||||
)
|
||||
stack.enter_context(patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False))
|
||||
stack.enter_context(patch.object(proxy_server, "llm_router", router))
|
||||
stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging))
|
||||
stack.enter_context(patch.object(proxy_server, "general_settings", {}))
|
||||
stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock()))
|
||||
stack.enter_context(patch.object(proxy_server, "version", "test-version"))
|
||||
stack.enter_context(patch.object(proxy_server, "prisma_client", MagicMock()))
|
||||
stack.enter_context(
|
||||
patch.object(proxy_server, "add_litellm_data_to_request", add_data)
|
||||
)
|
||||
stack.enter_context(patch.object(proxy_server, "add_litellm_data_to_request", add_data))
|
||||
|
||||
yield CancelHarness(
|
||||
data=data_holder,
|
||||
|
|
@ -1765,9 +1842,7 @@ async def test_cancel__model_encoded_id(cancel_harness):
|
|||
}
|
||||
|
||||
# OUTPUT SHAPE - response id re-encoded with the DECODED model.
|
||||
assert resp.id == encode_file_id_with_model(
|
||||
"batch-provider-id", "azure/gpt-4o", id_type="batch"
|
||||
)
|
||||
assert resp.id == encode_file_id_with_model("batch-provider-id", "azure/gpt-4o", id_type="batch")
|
||||
|
||||
# write-back tagged as a cancel.
|
||||
assert cancel_harness.update_batch_in_db.call_count == 1
|
||||
|
|
@ -1786,9 +1861,7 @@ async def test_cancel__model_encoded_id_forwards_deployment_model(cancel_harness
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__model_encoded_beats_unified(cancel_harness):
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID
|
||||
):
|
||||
with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID):
|
||||
await call_cancel(cancel_harness, AZURE_BATCH_ID)
|
||||
|
||||
assert cancel_harness.litellm_acancel.call_count == 1
|
||||
|
|
@ -1804,9 +1877,7 @@ async def test_cancel__model_encoded_beats_unified(cancel_harness):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__unified_batch_id_routes_to_router(cancel_harness):
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID
|
||||
):
|
||||
with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID):
|
||||
resp = await call_cancel(cancel_harness, "batch-unified-blob")
|
||||
|
||||
# DISPATCH - router fired, litellm did not, no creds lookup.
|
||||
|
|
@ -1845,8 +1916,9 @@ async def test_cancel__unified_missing_model_id_400(cancel_harness):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__unified_no_router_500(cancel_harness):
|
||||
with patch.object(proxy_server, "llm_router", None), patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID
|
||||
with (
|
||||
patch.object(proxy_server, "llm_router", None),
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_cancel(cancel_harness, "batch-unified-blob")
|
||||
|
|
@ -1885,9 +1957,7 @@ async def test_cancel__fallback_provider_path_param(cancel_harness):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel__fallback_provider_from_data_body(cancel_harness):
|
||||
await call_cancel(
|
||||
cancel_harness, "batch-raw-xyz", data_extra={"custom_llm_provider": "bedrock"}
|
||||
)
|
||||
await call_cancel(cancel_harness, "batch-raw-xyz", data_extra={"custom_llm_provider": "bedrock"})
|
||||
|
||||
assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "bedrock"
|
||||
|
||||
|
|
@ -1954,10 +2024,7 @@ async def test_cancel__exception_calls_failure_hook(cancel_harness):
|
|||
|
||||
cancel_harness.logging.post_call_failure_hook.assert_called_once()
|
||||
assert (
|
||||
cancel_harness.logging.post_call_failure_hook.call_args.kwargs[
|
||||
"original_exception"
|
||||
].args[0]
|
||||
== "provider boom"
|
||||
cancel_harness.logging.post_call_failure_hook.call_args.kwargs["original_exception"].args[0] == "provider boom"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1979,9 +2046,10 @@ async def test_create__loadbalancing_no_router_500(harness):
|
|||
},
|
||||
)
|
||||
harness.is_known_model.return_value = True
|
||||
with patch.object(
|
||||
litellm, "enable_loadbalancing_on_batch_endpoints", True
|
||||
), patch.object(proxy_server, "llm_router", None):
|
||||
with (
|
||||
patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True),
|
||||
patch.object(proxy_server, "llm_router", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_create(harness)
|
||||
|
||||
|
|
@ -2000,12 +2068,10 @@ async def test_create__unified_no_router_500(harness):
|
|||
"completion_window": "24h",
|
||||
},
|
||||
)
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"
|
||||
), patch.object(
|
||||
endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]
|
||||
), patch.object(
|
||||
proxy_server, "llm_router", None
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
|
||||
patch.object(endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"]),
|
||||
patch.object(proxy_server, "llm_router", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_create(harness)
|
||||
|
|
@ -2015,9 +2081,10 @@ async def test_create__unified_no_router_500(harness):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__unified_no_router_500(retrieve_harness):
|
||||
with patch.object(
|
||||
endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID
|
||||
), patch.object(proxy_server, "llm_router", None):
|
||||
with (
|
||||
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID),
|
||||
patch.object(proxy_server, "llm_router", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc:
|
||||
await call_retrieve(retrieve_harness, "batch-unified-blob")
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -11,6 +11,9 @@ Tests cover:
|
|||
- /v1/compress non-2xx surfaces as httpx.HTTPStatusError (raise_for_status),
|
||||
not a status_code check on the returned response -- both are handled
|
||||
- unreachable_fallback="fail_open" forwards the request uncompressed instead of raising
|
||||
- tokens_saved is derived from tokens_before/tokens_after when the compression
|
||||
service omits it, passed through verbatim when present, and skipped (without
|
||||
breaking compression) when the token counts are not numeric
|
||||
- CCR: headroom_retrieve tool injected when compressed messages contain hashes
|
||||
- CCR: async_should_run_agentic_loop returns True when response has headroom_retrieve tool calls
|
||||
- CCR: async_build_agentic_loop_plan calls retrieve endpoint and builds follow-up messages
|
||||
|
|
@ -32,6 +35,9 @@ from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import (
|
|||
has_headroom_retrieve_tool,
|
||||
HEADROOM_RETRIEVE_TOOL_NAME,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.compression_savings import (
|
||||
extract_compression_saved_tokens,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
FAKE_API_BASE = "https://headroom.example.com"
|
||||
|
|
@ -114,6 +120,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 +147,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 +157,126 @@ 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)
|
||||
|
||||
|
||||
def _recorded_guardrail_response(request_data: dict) -> dict:
|
||||
entries = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(entries) == 1
|
||||
return entries[0]["guardrail_response"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_derives_tokens_saved_when_service_omits_it(
|
||||
guardrail: HeadroomGuardrail,
|
||||
):
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["A" * 5000],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
# _make_compress_response omits tokens_saved, matching the live service.
|
||||
mock_response = _make_compress_response(COMPRESSED_MESSAGES)
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
stats = _recorded_guardrail_response(request_data)
|
||||
assert stats["tokens_saved"] == 900
|
||||
|
||||
# Spend tracking reads the entry under the spend-log metadata key.
|
||||
entry = request_data["metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert extract_compression_saved_tokens({"guardrail_information": [entry]}) == 900
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_passes_through_service_sent_tokens_saved(
|
||||
guardrail: HeadroomGuardrail,
|
||||
):
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["A" * 5000],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
mock_response = _make_compress_response(COMPRESSED_MESSAGES)
|
||||
# Deliberately different from tokens_before - tokens_after (900): the
|
||||
# service-sent value must win over the derived one.
|
||||
mock_response.json.return_value["tokens_saved"] = 123
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert _recorded_guardrail_response(request_data)["tokens_saved"] == 123
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"tokens_before, tokens_after",
|
||||
[
|
||||
("1000", "100"),
|
||||
(True, False),
|
||||
(None, None),
|
||||
],
|
||||
)
|
||||
async def test_apply_guardrail_skips_derivation_for_non_numeric_token_counts(
|
||||
guardrail: HeadroomGuardrail,
|
||||
tokens_before,
|
||||
tokens_after,
|
||||
):
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["A" * 5000],
|
||||
structured_messages=ORIGINAL_MESSAGES,
|
||||
)
|
||||
mock_response = _make_compress_response(COMPRESSED_MESSAGES)
|
||||
mock_response.json.return_value["tokens_before"] = tokens_before
|
||||
mock_response.json.return_value["tokens_after"] = tokens_after
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert "tokens_saved" not in _recorded_guardrail_response(request_data)
|
||||
# Compression itself is unaffected by the skipped derivation.
|
||||
assert result.get("structured_messages") == COMPRESSED_MESSAGES
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present(
|
||||
|
|
@ -719,6 +858,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 +869,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 +888,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 +1058,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 +1068,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
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ imports these inside function bodies to avoid circular imports.
|
|||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -203,6 +203,71 @@ class TestToolManagementEndpoints:
|
|||
assert tuple(call.args[1:]) == expected_binds
|
||||
assert resp.json()["end_date"] == "2026-07-02"
|
||||
|
||||
def test_tool_spend_start_clamped_to_30_days_before_end(self):
|
||||
# Clamped floor is end_date minus 30 days, serving up to 31 calendar dates
|
||||
# inclusive: deliberately the same width as the endpoint's default window,
|
||||
# so the dashboard's default range never triggers the clamp.
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(return_value=[])
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
resp = self.client.get("/v1/tool/spend?start_date=2026-01-01&end_date=2026-07-01")
|
||||
assert resp.status_code == 200
|
||||
expected_binds = (
|
||||
datetime(2026, 6, 1, tzinfo=timezone.utc).isoformat(),
|
||||
datetime(2026, 7, 2, tzinfo=timezone.utc).isoformat(),
|
||||
)
|
||||
assert prisma.db.query_raw.await_count == 2
|
||||
for call in prisma.db.query_raw.await_args_list:
|
||||
assert tuple(call.args[1:]) == expected_binds
|
||||
assert resp.json()["start_date"] == "2026-06-01"
|
||||
assert resp.json()["end_date"] == "2026-07-01"
|
||||
|
||||
def test_tool_spend_range_within_cap_is_not_clamped(self):
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(return_value=[])
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
resp = self.client.get("/v1/tool/spend?start_date=2026-06-25&end_date=2026-07-01")
|
||||
assert resp.status_code == 200
|
||||
for call in prisma.db.query_raw.await_args_list:
|
||||
assert call.args[1] == datetime(2026, 6, 25, tzinfo=timezone.utc).isoformat()
|
||||
assert resp.json()["start_date"] == "2026-06-25"
|
||||
|
||||
def test_tool_spend_start_honored_when_end_date_omitted(self):
|
||||
# Regression: with end_date omitted the floor anchors to today's UTC
|
||||
# midnight, not now's time-of-day, so an explicit start_date exactly 30
|
||||
# days back is served from midnight rather than truncated to mid-day.
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(return_value=[])
|
||||
floor_day = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) - timedelta(days=30)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
resp = self.client.get(f"/v1/tool/spend?start_date={floor_day.strftime('%Y-%m-%d')}")
|
||||
assert resp.status_code == 200
|
||||
for call in prisma.db.query_raw.await_args_list:
|
||||
assert call.args[1] == floor_day.isoformat()
|
||||
assert resp.json()["start_date"] == floor_day.strftime("%Y-%m-%d")
|
||||
|
||||
def test_tool_spend_clamp_without_end_date_lands_on_midnight(self):
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(return_value=[])
|
||||
floor_day = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0) - timedelta(days=30)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
resp = self.client.get("/v1/tool/spend?start_date=2020-01-01")
|
||||
assert resp.status_code == 200
|
||||
for call in prisma.db.query_raw.await_args_list:
|
||||
assert call.args[1] == floor_day.isoformat()
|
||||
assert resp.json()["start_date"] == floor_day.strftime("%Y-%m-%d")
|
||||
|
||||
def test_tool_spend_total_query_bounds_outer_spendlogs_scan(self):
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(return_value=[])
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02")
|
||||
assert resp.status_code == 200
|
||||
for call in prisma.db.query_raw.await_args_list:
|
||||
sql = call.args[0]
|
||||
assert 'sl."startTime" >=' in sql
|
||||
assert 'sl."startTime" <' in sql
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"query",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -157,8 +157,9 @@ async def test_cleanup_old_spend_logs_batch_deletion():
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = MagicMock()
|
||||
|
||||
# Mock execute_raw to return deleted counts
|
||||
mock_db.execute_raw = AsyncMock(side_effect=[1000, 500, 0])
|
||||
# Mock execute_raw to return deleted counts (3 spend-log batches, then the
|
||||
# tool-index cleanup's first batch returning 0)
|
||||
mock_db.execute_raw = AsyncMock(side_effect=[1000, 500, 0, 0])
|
||||
|
||||
# Wire up mocks
|
||||
mock_prisma_client.db = mock_db
|
||||
|
|
@ -178,7 +179,7 @@ async def test_cleanup_old_spend_logs_batch_deletion():
|
|||
await cleaner.cleanup_old_spend_logs(mock_prisma_client)
|
||||
|
||||
# Validate batching and deletion via raw SQL
|
||||
assert mock_db.execute_raw.call_count == 3
|
||||
assert mock_db.execute_raw.call_count == 4
|
||||
|
||||
# Check the first call argument
|
||||
call_args_sql = mock_db.execute_raw.call_args_list[0][0][0]
|
||||
|
|
@ -188,6 +189,10 @@ async def test_cleanup_old_spend_logs_batch_deletion():
|
|||
# reusing x-litellm-call-id take out a fresh row alongside the expired one
|
||||
assert 'WHERE ("request_id", "startTime") IN' in call_args_sql
|
||||
|
||||
# After spend logs, the derived tool index rows expire on the same cutoff
|
||||
tool_index_sql = mock_db.execute_raw.call_args_list[3][0][0]
|
||||
assert 'DELETE FROM "LiteLLM_SpendLogToolIndex"' in tool_index_sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleanup_old_spend_logs_retention_period_cutoff():
|
||||
|
|
@ -258,6 +263,10 @@ async def test_cleanup_drops_partitions_when_enabled_and_partitioned():
|
|||
partition_manager.drop_partitions_older_than.assert_awaited_once()
|
||||
delete_sql = mock_prisma_client.db.execute_raw.call_args_list[0][0][0]
|
||||
assert 'DELETE FROM "LiteLLM_SpendLogs"' in delete_sql
|
||||
# Partition drops only reclaim spend logs; the tool index must still be
|
||||
# cleaned row-wise on the same run
|
||||
all_sql = [c[0][0] for c in mock_prisma_client.db.execute_raw.call_args_list]
|
||||
assert any('DELETE FROM "LiteLLM_SpendLogToolIndex"' in s for s in all_sql)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -270,7 +279,7 @@ async def test_cleanup_uses_delete_when_partitioning_not_enabled():
|
|||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[10, 0])
|
||||
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[10, 0, 0])
|
||||
|
||||
partition_manager = MagicMock()
|
||||
partition_manager.is_partitioned = AsyncMock(return_value=True)
|
||||
|
|
@ -301,7 +310,7 @@ async def test_cleanup_uses_delete_when_not_partitioned():
|
|||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[10, 0])
|
||||
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[10, 0, 0])
|
||||
|
||||
partition_manager = MagicMock()
|
||||
partition_manager.is_partitioned = AsyncMock(return_value=False)
|
||||
|
|
@ -320,7 +329,7 @@ async def test_cleanup_uses_delete_when_not_partitioned():
|
|||
await cleaner.cleanup_old_spend_logs(mock_prisma_client)
|
||||
|
||||
partition_manager.drop_partitions_older_than.assert_not_awaited()
|
||||
assert mock_prisma_client.db.execute_raw.await_count == 2
|
||||
assert mock_prisma_client.db.execute_raw.await_count == 3
|
||||
delete_sql = mock_prisma_client.db.execute_raw.call_args_list[0][0][0]
|
||||
assert 'DELETE FROM "LiteLLM_SpendLogs"' in delete_sql
|
||||
|
||||
|
|
@ -437,6 +446,55 @@ async def test_delete_old_logs_continues_on_valid_int_return():
|
|||
assert total_deleted == 800
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_old_rows_stops_at_max_batches(monkeypatch):
|
||||
"""The run-loop backstop must halt a cleanup that keeps finding rows, so a
|
||||
huge backlog is spread across scheduled runs instead of one unbounded loop."""
|
||||
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
|
||||
|
||||
monkeypatch.setattr(cleanup_module, "SPEND_LOG_RUN_LOOPS", 2)
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_db = MagicMock()
|
||||
mock_db.execute_raw = AsyncMock(return_value=1000)
|
||||
mock_prisma_client.db = mock_db
|
||||
|
||||
cleaner = SpendLogCleanup(
|
||||
general_settings={"maximum_spend_logs_retention_period": "7d"}
|
||||
)
|
||||
|
||||
cutoff_date = datetime.now(timezone.utc) - timedelta(days=7)
|
||||
total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date)
|
||||
|
||||
# run_count exceeds the cap only after 3 full batches (0, 1, 2)
|
||||
assert mock_db.execute_raw.call_count == 3
|
||||
assert total_deleted == 3000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_old_tool_index_rows_deletes_on_composite_key():
|
||||
"""Tool index rows are derived from spend logs and expire on the same cutoff;
|
||||
the delete must match on the table's composite primary key."""
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_db = MagicMock()
|
||||
mock_db.execute_raw = AsyncMock(side_effect=[5, 0])
|
||||
mock_prisma_client.db = mock_db
|
||||
|
||||
cleaner = SpendLogCleanup(
|
||||
general_settings={"maximum_spend_logs_retention_period": "7d"}
|
||||
)
|
||||
|
||||
cutoff_date = datetime.now(timezone.utc) - timedelta(days=7)
|
||||
total_deleted = await cleaner._delete_old_tool_index_rows(mock_prisma_client, cutoff_date)
|
||||
|
||||
assert total_deleted == 5
|
||||
delete_sql = mock_db.execute_raw.call_args_list[0][0][0]
|
||||
assert 'DELETE FROM "LiteLLM_SpendLogToolIndex"' in delete_sql
|
||||
assert 'WHERE ("request_id", "tool_name") IN' in delete_sql
|
||||
assert '"start_time" <' in delete_sql
|
||||
assert mock_db.execute_raw.call_args_list[0][0][1] == cutoff_date
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_old_logs_continues_after_single_batch_failure(monkeypatch):
|
||||
"""A single batch failure (e.g. DB timeout) must not abort the whole run —
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -2964,11 +2959,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/common_components/TableHeaderSortDropdown/TableHeaderSortDropdown.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/common_components/budget_duration_dropdown.tsx": {
|
||||
"local/filename-pascal-case": {
|
||||
"count": 1
|
||||
|
|
@ -3554,11 +3544,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/routing_groups/RoutingGroupsTable.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 2
|
||||
}
|
||||
},
|
||||
"src/components/routing_groups/index.tsx": {
|
||||
"local/filename-pascal-case": {
|
||||
"count": 1
|
||||
|
|
|
|||
|
|
@ -178,4 +178,31 @@ describe("UsageTab", () => {
|
|||
(await findByTestId("bar-0-my tool/read")).click();
|
||||
expect(mockPush).toHaveBeenLastCalledWith("/ui/tool-policies?tool=my+tool%2Fread");
|
||||
});
|
||||
|
||||
it("notes the 30-day cap when the server clamps the tool spend window", async () => {
|
||||
const toolSpend = {
|
||||
by_tool: [{ tool_name: "search", spend: 4.0, call_count: 3, total_tokens: 150 }],
|
||||
daily: [{ date: "2026-07-12", tool_name: "search", spend: 4.0, call_count: 3 }],
|
||||
total_spend: 4.0,
|
||||
start_date: "2026-07-05",
|
||||
end_date: "2026-07-14",
|
||||
};
|
||||
const { findByText } = renderWith([day("2026-07-12", {})], toolSpend);
|
||||
|
||||
expect(await findByText(/capped at 30 days before the end of the selected range/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows no cap note when the served window matches the request", async () => {
|
||||
const toolSpend = {
|
||||
by_tool: [{ tool_name: "search", spend: 4.0, call_count: 3, total_tokens: 150 }],
|
||||
daily: [{ date: "2026-07-12", tool_name: "search", spend: 4.0, call_count: 3 }],
|
||||
total_spend: 4.0,
|
||||
start_date: "2026-07-01",
|
||||
end_date: "2026-07-14",
|
||||
};
|
||||
const { findAllByTestId, queryByText } = renderWith([day("2026-07-12", {})], toolSpend);
|
||||
|
||||
await findAllByTestId("bar-chart");
|
||||
expect(queryByText(/capped at 30 days before the end of the selected range/)).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -120,6 +120,7 @@ const UsageTab: React.FC<UsageTabProps> = ({ accessToken, activity }) => {
|
|||
|
||||
const toolSpend = toolSpendState?.key === rangeKey ? toolSpendState.data : null;
|
||||
const toolSpendLoading = toolSpendEnabled && toolSpend === null;
|
||||
const toolSpendWindowClamped = !!toolSpend?.start_date && !!startTime && toolSpend.start_date > isoDay(startTime);
|
||||
|
||||
const compressionTotal = useMemo(() => results.reduce((sum, d) => sum + compressionOf(d.metrics), 0), [results]);
|
||||
const cachingTotal = useMemo(() => results.reduce((sum, d) => sum + cachingOf(d.metrics), 0), [results]);
|
||||
|
|
@ -224,6 +225,12 @@ const UsageTab: React.FC<UsageTabProps> = ({ accessToken, activity }) => {
|
|||
counts its full spend toward each, so this attributes rather than partitions spend. Click a bar to see the
|
||||
logs for that tool.
|
||||
</p>
|
||||
{toolSpendWindowClamped && (
|
||||
<p className="text-xs text-muted-foreground">
|
||||
Tool spend is capped at 30 days before the end of the selected range; showing spend since{" "}
|
||||
{toolSpend?.start_date}.
|
||||
</p>
|
||||
)}
|
||||
</CardHeader>
|
||||
<CardContent>
|
||||
{topTools.length === 0 ? (
|
||||
|
|
|
|||
|
|
@ -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,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,148 +0,0 @@
|
|||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { TableHeaderSortDropdown } from "./TableHeaderSortDropdown";
|
||||
|
||||
describe("TableHeaderSortDropdown", () => {
|
||||
it("should render", () => {
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState={false} onSortChange={onSortChange} />);
|
||||
expect(screen.getByRole("button")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should open dropdown menu when button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState={false} onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Ascending")).toBeInTheDocument();
|
||||
expect(screen.getByText("Descending")).toBeInTheDocument();
|
||||
expect(screen.getByText("Reset")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should call onSortChange with asc when ascending option is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState={false} onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Ascending")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const ascendingOption = screen.getByText("Ascending");
|
||||
await user.click(ascendingOption);
|
||||
|
||||
expect(onSortChange).toHaveBeenCalledTimes(1);
|
||||
expect(onSortChange).toHaveBeenCalledWith("asc");
|
||||
});
|
||||
|
||||
it("should call onSortChange with desc when descending option is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState={false} onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Descending")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const descendingOption = screen.getByText("Descending");
|
||||
await user.click(descendingOption);
|
||||
|
||||
expect(onSortChange).toHaveBeenCalledTimes(1);
|
||||
expect(onSortChange).toHaveBeenCalledWith("desc");
|
||||
});
|
||||
|
||||
it("should call onSortChange with false when reset option is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState="asc" onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Reset")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const resetOption = screen.getByText("Reset");
|
||||
await user.click(resetOption);
|
||||
|
||||
expect(onSortChange).toHaveBeenCalledTimes(1);
|
||||
expect(onSortChange).toHaveBeenCalledWith(false);
|
||||
});
|
||||
|
||||
it("should highlight ascending option when sort state is asc", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState="asc" onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
const ascendingOption = screen.getByText("Ascending");
|
||||
const menuItem = ascendingOption.closest(".ant-dropdown-menu-item");
|
||||
expect(menuItem).toHaveClass("ant-dropdown-menu-item-selected");
|
||||
});
|
||||
});
|
||||
|
||||
it("should highlight descending option when sort state is desc", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState="desc" onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
const descendingOption = screen.getByText("Descending");
|
||||
const menuItem = descendingOption.closest(".ant-dropdown-menu-item");
|
||||
expect(menuItem).toHaveClass("ant-dropdown-menu-item-selected");
|
||||
});
|
||||
});
|
||||
|
||||
it("should not highlight any option when sort state is false", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
render(<TableHeaderSortDropdown sortState={false} onSortChange={onSortChange} />);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Ascending")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const ascendingOption = screen.getByText("Ascending");
|
||||
const menuItem = ascendingOption.closest(".ant-dropdown-menu-item");
|
||||
expect(menuItem).not.toHaveClass("ant-dropdown-menu-item-selected");
|
||||
});
|
||||
|
||||
it("should stop event propagation when button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onSortChange = vi.fn();
|
||||
const onParentClick = vi.fn();
|
||||
|
||||
render(
|
||||
<div onClick={onParentClick}>
|
||||
<TableHeaderSortDropdown sortState={false} onSortChange={onSortChange} />
|
||||
</div>,
|
||||
);
|
||||
|
||||
const button = screen.getByRole("button");
|
||||
await user.click(button);
|
||||
|
||||
expect(onParentClick).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,82 +0,0 @@
|
|||
import React from "react";
|
||||
import { Button, Dropdown, MenuProps } from "antd";
|
||||
import { SwitchVerticalIcon, ChevronUpIcon, ChevronDownIcon, XIcon } from "@heroicons/react/outline";
|
||||
|
||||
export type SortState = "asc" | "desc" | false;
|
||||
|
||||
interface TableHeaderSortDropdownProps {
|
||||
/**
|
||||
* Current sort state: "asc", "desc", or false for neutral
|
||||
*/
|
||||
sortState: SortState;
|
||||
/**
|
||||
* Callback when sort state changes
|
||||
* @param newState - The new sort state: "asc", "desc", or false
|
||||
*/
|
||||
onSortChange: (newState: SortState) => void;
|
||||
/**
|
||||
* Optional column ID for identification
|
||||
*/
|
||||
columnId?: string;
|
||||
}
|
||||
|
||||
export const TableHeaderSortDropdown: React.FC<TableHeaderSortDropdownProps> = ({ sortState, onSortChange }) => {
|
||||
const handleMenuClick: MenuProps["onClick"] = ({ key }) => {
|
||||
if (key === "asc") {
|
||||
onSortChange("asc");
|
||||
} else if (key === "desc") {
|
||||
onSortChange("desc");
|
||||
} else if (key === "reset") {
|
||||
onSortChange(false);
|
||||
}
|
||||
};
|
||||
|
||||
const menuItems: MenuProps["items"] = [
|
||||
{
|
||||
key: "asc",
|
||||
label: "Ascending",
|
||||
icon: <ChevronUpIcon className="h-4 w-4" />,
|
||||
},
|
||||
{
|
||||
key: "desc",
|
||||
label: "Descending",
|
||||
icon: <ChevronDownIcon className="h-4 w-4" />,
|
||||
},
|
||||
{
|
||||
key: "reset",
|
||||
label: "Reset",
|
||||
icon: <XIcon className="h-4 w-4" />,
|
||||
},
|
||||
];
|
||||
|
||||
// Determine which icon to display based on current sort state
|
||||
const renderIcon = () => {
|
||||
if (sortState === "asc") {
|
||||
return <ChevronUpIcon className="h-4 w-4" />;
|
||||
} else if (sortState === "desc") {
|
||||
return <ChevronDownIcon className="h-4 w-4" />;
|
||||
} else {
|
||||
return <SwitchVerticalIcon className="h-4 w-4" />;
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<Dropdown
|
||||
menu={{
|
||||
items: menuItems,
|
||||
onClick: handleMenuClick,
|
||||
selectable: true,
|
||||
selectedKeys: sortState ? [sortState] : [],
|
||||
}}
|
||||
trigger={["click"]}
|
||||
autoAdjustOverflow
|
||||
>
|
||||
<Button
|
||||
type="text"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
icon={renderIcon()}
|
||||
className={sortState ? "text-blue-500 hover:text-blue-600" : "text-gray-400 hover:text-blue-500"}
|
||||
/>
|
||||
</Dropdown>
|
||||
);
|
||||
};
|
||||
|
|
@ -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)}
|
||||
>
|
||||
|
|
|
|||
|
|
@ -0,0 +1,91 @@
|
|||
"use client";
|
||||
|
||||
import { Code2 } from "lucide-react";
|
||||
import React from "react";
|
||||
|
||||
import CodeBlock from "@/components/CodeBlock";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
|
||||
import { formatStrategyLabel } from "./strategy";
|
||||
import type { RoutingGroup } from "./types";
|
||||
|
||||
interface RoutingGroupUsagePanelProps {
|
||||
group: RoutingGroup;
|
||||
baseUrl: string;
|
||||
}
|
||||
|
||||
const exampleModel = (group: RoutingGroup): string => group.models[0] ?? "<your-model>";
|
||||
|
||||
const buildCurlSnippet = (group: RoutingGroup, baseUrl: string): string =>
|
||||
`curl -X POST '${baseUrl}/v1/chat/completions' \\
|
||||
-H 'Content-Type: application/json' \\
|
||||
-H 'Authorization: Bearer $LITELLM_API_KEY' \\
|
||||
-d '{
|
||||
"model": "${exampleModel(group)}",
|
||||
"messages": [{"role": "user", "content": "Hello!"}]
|
||||
}'`;
|
||||
|
||||
const buildPythonSnippet = (group: RoutingGroup, baseUrl: string): string =>
|
||||
`from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
api_key="$LITELLM_API_KEY",
|
||||
base_url="${baseUrl}",
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="${exampleModel(group)}",
|
||||
messages=[{"role": "user", "content": "Hello!"}],
|
||||
)
|
||||
|
||||
print(response)`;
|
||||
|
||||
const buildJsSnippet = (group: RoutingGroup, baseUrl: string): string =>
|
||||
`import OpenAI from "openai";
|
||||
|
||||
const client = new OpenAI({
|
||||
apiKey: process.env.LITELLM_API_KEY,
|
||||
baseURL: "${baseUrl}",
|
||||
});
|
||||
|
||||
const response = await client.chat.completions.create({
|
||||
model: "${exampleModel(group)}",
|
||||
messages: [{ role: "user", content: "Hello!" }],
|
||||
});
|
||||
|
||||
console.log(response);`;
|
||||
|
||||
const SNIPPET_TABS = [
|
||||
{ value: "curl", label: "cURL", language: "bash", build: buildCurlSnippet },
|
||||
{ value: "python", label: "Python (OpenAI SDK)", language: "python", build: buildPythonSnippet },
|
||||
{ value: "javascript", label: "JavaScript (OpenAI SDK)", language: "javascript", build: buildJsSnippet },
|
||||
] as const;
|
||||
|
||||
export function RoutingGroupUsagePanel({ group, baseUrl }: RoutingGroupUsagePanelProps) {
|
||||
return (
|
||||
<div className="border-y bg-muted/40 px-4 py-4">
|
||||
<div className="mb-2 flex items-center gap-2">
|
||||
<Code2 className="size-4 text-primary" />
|
||||
<span className="text-sm font-medium text-foreground">How routing works for this group</span>
|
||||
</div>
|
||||
<p className="mb-3 text-sm text-muted-foreground">
|
||||
Callers request any model in the group by name; LiteLLM picks a deployment behind the scenes using the{" "}
|
||||
<span className="font-medium text-foreground">{formatStrategyLabel(group.routing_strategy)}</span> strategy.
|
||||
</p>
|
||||
<Tabs defaultValue="curl">
|
||||
<TabsList variant="line" className="h-auto w-full justify-start rounded-none border-b p-0">
|
||||
{SNIPPET_TABS.map((tab) => (
|
||||
<TabsTrigger key={tab.value} value={tab.value} className="flex-none rounded-none px-4 py-2">
|
||||
{tab.label}
|
||||
</TabsTrigger>
|
||||
))}
|
||||
</TabsList>
|
||||
{SNIPPET_TABS.map((tab) => (
|
||||
<TabsContent key={tab.value} value={tab.value} className="pt-3">
|
||||
<CodeBlock language={tab.language} code={tab.build(group, baseUrl)} />
|
||||
</TabsContent>
|
||||
))}
|
||||
</Tabs>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,145 @@
|
|||
import { render, screen, within } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
import RoutingGroupsTable from "./RoutingGroupsTable";
|
||||
import type { RoutingGroup } from "./types";
|
||||
|
||||
describe("RoutingGroupsTable", () => {
|
||||
const onEdit = vi.fn();
|
||||
const onDelete = vi.fn();
|
||||
|
||||
const prodGroup: RoutingGroup = {
|
||||
group_name: "prod-group",
|
||||
models: ["gpt-4o", "claude-sonnet-4-5"],
|
||||
routing_strategy: "usage-based-routing",
|
||||
};
|
||||
|
||||
const devGroup: RoutingGroup = {
|
||||
group_name: "dev-group",
|
||||
models: ["gpt-4o-mini"],
|
||||
routing_strategy: "simple-shuffle",
|
||||
};
|
||||
|
||||
const defaultProps = {
|
||||
groups: [] as RoutingGroup[],
|
||||
onEdit,
|
||||
onDelete,
|
||||
proxyBaseUrl: "https://proxy.example.com",
|
||||
};
|
||||
|
||||
const rowFor = (groupName: string): HTMLElement => {
|
||||
const row = document.querySelector(`[data-row-id="${groupName}"]`);
|
||||
if (!(row instanceof HTMLElement)) {
|
||||
throw new Error(`No row rendered for ${groupName}`);
|
||||
}
|
||||
return row;
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render every column header", () => {
|
||||
render(<RoutingGroupsTable {...defaultProps} />);
|
||||
for (const header of ["Group Name", "Models", "Strategy"]) {
|
||||
expect(screen.getByText(header)).toBeInTheDocument();
|
||||
}
|
||||
});
|
||||
|
||||
it("should show the empty state when there are no groups", () => {
|
||||
render(<RoutingGroupsTable {...defaultProps} />);
|
||||
expect(screen.getByText("No routing groups yet")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the group name, its models, and a human-readable strategy label", () => {
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup]} />);
|
||||
const row = rowFor("prod-group");
|
||||
expect(within(row).getByText("prod-group")).toBeInTheDocument();
|
||||
expect(within(row).getByText("gpt-4o")).toBeInTheDocument();
|
||||
expect(within(row).getByText("claude-sonnet-4-5")).toBeInTheDocument();
|
||||
expect(within(row).getByText("Usage Based")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should fall back to the raw strategy value when it has no friendly label", () => {
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[{ ...prodGroup, routing_strategy: "custom-strategy" }]} />);
|
||||
expect(within(rowFor("prod-group")).getByText("custom-strategy")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should collapse models beyond the first three behind a +N more badge", () => {
|
||||
const wideGroup: RoutingGroup = { ...prodGroup, models: ["a", "b", "c", "d", "e"] };
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[wideGroup]} />);
|
||||
const row = rowFor("prod-group");
|
||||
expect(within(row).getByText("+2 more")).toBeInTheDocument();
|
||||
expect(within(row).queryByText("d")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should keep the incoming order until a column is sorted", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup, devGroup]} />);
|
||||
|
||||
const namesInOrder = () =>
|
||||
screen
|
||||
.getAllByRole("row")
|
||||
.slice(1)
|
||||
.map((row) => row.getAttribute("data-row-id"));
|
||||
|
||||
expect(namesInOrder()).toEqual(["prod-group", "dev-group"]);
|
||||
|
||||
await user.click(screen.getByTestId("sort-header-group_name"));
|
||||
expect(namesInOrder()).toEqual(["dev-group", "prod-group"]);
|
||||
});
|
||||
|
||||
it("should toggle the usage panel when the group name is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup]} />);
|
||||
|
||||
expect(screen.queryByText("How routing works for this group")).not.toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "prod-group" }));
|
||||
expect(await screen.findByText("How routing works for this group")).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "prod-group" }));
|
||||
expect(screen.queryByText("How routing works for this group")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should build the usage snippet from the proxy base url and the group's first model", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup]} />);
|
||||
await user.click(screen.getByRole("button", { name: "prod-group" }));
|
||||
|
||||
const panel = (await screen.findByText("How routing works for this group")).closest("div")?.parentElement;
|
||||
expect(panel?.textContent).toContain("https://proxy.example.com");
|
||||
expect(panel?.textContent).toContain("gpt-4o");
|
||||
});
|
||||
|
||||
it("should expand only the clicked group", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup, devGroup]} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "dev-group" }));
|
||||
expect(await screen.findAllByText("How routing works for this group")).toHaveLength(1);
|
||||
expect(within(rowFor("prod-group")).queryByText("How routing works for this group")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should edit a group through the actions menu", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup]} />);
|
||||
await user.click(screen.getByTestId("routing-group-actions-prod-group"));
|
||||
await user.click(await screen.findByTestId("routing-group-action-edit"));
|
||||
expect(onEdit).toHaveBeenCalledWith(prodGroup);
|
||||
});
|
||||
|
||||
it("should delete a group through the actions menu", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<RoutingGroupsTable {...defaultProps} groups={[prodGroup]} />);
|
||||
await user.click(screen.getByTestId("routing-group-actions-prod-group"));
|
||||
await user.click(await screen.findByTestId("routing-group-action-delete"));
|
||||
expect(onDelete).toHaveBeenCalledWith(prodGroup);
|
||||
});
|
||||
|
||||
it("should show skeleton rows instead of the empty state while loading", () => {
|
||||
render(<RoutingGroupsTable {...defaultProps} isLoading />);
|
||||
expect(screen.queryByText("No routing groups yet")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,229 +1,82 @@
|
|||
"use client";
|
||||
|
||||
import React, { useState } from "react";
|
||||
import { Flex, Table, Tabs, Tag, Tooltip, Typography, Button } from "antd";
|
||||
import type { ColumnsType } from "antd/es/table";
|
||||
import { BranchesOutlined, DeleteOutlined, EditOutlined, CodeOutlined } from "@ant-design/icons";
|
||||
import type { RoutingGroup } from "./types";
|
||||
import type { ExpandedState, SortingState } from "@tanstack/react-table";
|
||||
import { Inbox } from "lucide-react";
|
||||
import React, { useCallback, useMemo, useState } from "react";
|
||||
|
||||
const { Text, Paragraph } = Typography;
|
||||
import { DataTable } from "@/components/shared/DataTable";
|
||||
|
||||
import { RoutingGroupUsagePanel } from "./RoutingGroupUsagePanel";
|
||||
import { getRoutingGroupsTableColumns } from "./RoutingGroupsTableColumns";
|
||||
import type { RoutingGroup } from "./types";
|
||||
|
||||
interface RoutingGroupsTableProps {
|
||||
groups: RoutingGroup[];
|
||||
loading?: boolean;
|
||||
isLoading?: boolean;
|
||||
onEdit: (group: RoutingGroup) => void;
|
||||
onDelete: (group: RoutingGroup) => void;
|
||||
proxyBaseUrl?: string;
|
||||
}
|
||||
|
||||
const formatStrategyLabel = (strategy: string): string => {
|
||||
switch (strategy) {
|
||||
case "simple-shuffle":
|
||||
return "Simple Shuffle";
|
||||
case "least-busy":
|
||||
return "Least Busy";
|
||||
case "usage-based-routing":
|
||||
return "Usage Based";
|
||||
case "latency-based-routing":
|
||||
return "Latency Based";
|
||||
default:
|
||||
return strategy;
|
||||
}
|
||||
};
|
||||
|
||||
const resolveBaseUrl = (proxyBaseUrl?: string): string => {
|
||||
if (proxyBaseUrl && proxyBaseUrl.trim()) return proxyBaseUrl;
|
||||
if (typeof window !== "undefined" && window.location?.origin) return window.location.origin;
|
||||
return "<your_proxy_base_url>";
|
||||
};
|
||||
|
||||
const exampleModel = (group: RoutingGroup): string => group.models[0] ?? "<your-model>";
|
||||
|
||||
const buildCurlSnippet = (group: RoutingGroup, baseUrl: string): string =>
|
||||
`curl -X POST '${baseUrl}/v1/chat/completions' \\
|
||||
-H 'Content-Type: application/json' \\
|
||||
-H 'Authorization: Bearer $LITELLM_API_KEY' \\
|
||||
-d '{
|
||||
"model": "${exampleModel(group)}",
|
||||
"messages": [{"role": "user", "content": "Hello!"}]
|
||||
}'`;
|
||||
|
||||
const buildPythonSnippet = (group: RoutingGroup, baseUrl: string): string =>
|
||||
`from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
api_key="$LITELLM_API_KEY",
|
||||
base_url="${baseUrl}",
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="${exampleModel(group)}",
|
||||
messages=[{"role": "user", "content": "Hello!"}],
|
||||
)
|
||||
|
||||
print(response)`;
|
||||
|
||||
const buildJsSnippet = (group: RoutingGroup, baseUrl: string): string =>
|
||||
`import OpenAI from "openai";
|
||||
|
||||
const client = new OpenAI({
|
||||
apiKey: process.env.LITELLM_API_KEY,
|
||||
baseURL: "${baseUrl}",
|
||||
});
|
||||
|
||||
const response = await client.chat.completions.create({
|
||||
model: "${exampleModel(group)}",
|
||||
messages: [{ role: "user", content: "Hello!" }],
|
||||
});
|
||||
|
||||
console.log(response);`;
|
||||
|
||||
interface RoutingGroupSnippetProps {
|
||||
group: RoutingGroup;
|
||||
baseUrl: string;
|
||||
function EmptyState() {
|
||||
return (
|
||||
<div className="flex flex-col items-center gap-1 py-6">
|
||||
<div className="mb-1 flex size-10 items-center justify-center rounded-lg bg-muted">
|
||||
<Inbox className="size-5 text-muted-foreground" />
|
||||
</div>
|
||||
<div className="text-sm font-medium text-foreground">No routing groups yet</div>
|
||||
<div className="text-sm text-muted-foreground">
|
||||
Create a group to load-balance a set of models behind one name.
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
const SNIPPET_BLOCK_STYLE: React.CSSProperties = {
|
||||
backgroundColor: "#111827",
|
||||
color: "#f3f4f6",
|
||||
borderRadius: 6,
|
||||
padding: 16,
|
||||
fontSize: 12,
|
||||
whiteSpace: "pre",
|
||||
overflowX: "auto",
|
||||
};
|
||||
|
||||
const RoutingGroupSnippet: React.FC<RoutingGroupSnippetProps> = ({ group, baseUrl }) => {
|
||||
const snippets = {
|
||||
curl: buildCurlSnippet(group, baseUrl),
|
||||
python: buildPythonSnippet(group, baseUrl),
|
||||
javascript: buildJsSnippet(group, baseUrl),
|
||||
} as const;
|
||||
type SnippetKey = keyof typeof snippets;
|
||||
const [activeKey, setActiveKey] = useState<SnippetKey>("curl");
|
||||
|
||||
const items = [
|
||||
{ key: "curl", label: "cURL" },
|
||||
{ key: "python", label: "Python (OpenAI SDK)" },
|
||||
{ key: "javascript", label: "JavaScript (OpenAI SDK)" },
|
||||
].map(({ key, label }) => ({
|
||||
key,
|
||||
label,
|
||||
children: (
|
||||
<Paragraph code className="mb-0!" style={SNIPPET_BLOCK_STYLE}>
|
||||
{snippets[key as SnippetKey]}
|
||||
</Paragraph>
|
||||
),
|
||||
}));
|
||||
|
||||
return (
|
||||
<Tabs
|
||||
size="small"
|
||||
activeKey={activeKey}
|
||||
onChange={(k) => setActiveKey(k as SnippetKey)}
|
||||
items={items}
|
||||
tabBarExtraContent={
|
||||
<Paragraph copyable={{ text: snippets[activeKey], tooltips: ["Copy", "Copied"] }} className="mb-0!" />
|
||||
}
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
||||
const RoutingGroupsTable: React.FC<RoutingGroupsTableProps> = ({ groups, loading, onEdit, onDelete, proxyBaseUrl }) => {
|
||||
const [expandedRowKeys, setExpandedRowKeys] = useState<React.Key[]>([]);
|
||||
const RoutingGroupsTable: React.FC<RoutingGroupsTableProps> = ({
|
||||
groups,
|
||||
isLoading,
|
||||
onEdit,
|
||||
onDelete,
|
||||
proxyBaseUrl,
|
||||
}) => {
|
||||
const [sorting, setSorting] = useState<SortingState>([]);
|
||||
const [expanded, setExpanded] = useState<ExpandedState>({});
|
||||
const baseUrl = resolveBaseUrl(proxyBaseUrl);
|
||||
|
||||
const columns: ColumnsType<RoutingGroup> = [
|
||||
{
|
||||
title: "GROUP NAME",
|
||||
dataIndex: "group_name",
|
||||
key: "group_name",
|
||||
render: (name: string) => (
|
||||
<Text strong className="text-blue-600">
|
||||
{name}
|
||||
</Text>
|
||||
),
|
||||
},
|
||||
{
|
||||
title: "MODELS",
|
||||
dataIndex: "models",
|
||||
key: "models",
|
||||
render: (models: string[]) => (
|
||||
<Flex wrap="wrap" gap={4}>
|
||||
{models.map((m) => (
|
||||
<Tag key={m}>{m}</Tag>
|
||||
))}
|
||||
</Flex>
|
||||
),
|
||||
},
|
||||
{
|
||||
title: "STRATEGY",
|
||||
dataIndex: "routing_strategy",
|
||||
key: "routing_strategy",
|
||||
render: (strategy: string) => (
|
||||
<span className="inline-flex items-center gap-1.5">
|
||||
<BranchesOutlined className="text-gray-400" />
|
||||
<Text>{formatStrategyLabel(strategy)}</Text>
|
||||
</span>
|
||||
),
|
||||
},
|
||||
{
|
||||
title: "ACTIONS",
|
||||
key: "actions",
|
||||
width: 120,
|
||||
align: "right",
|
||||
render: (_, group) => (
|
||||
<Flex justify="flex-end" align="center" gap={8}>
|
||||
<Tooltip title="Edit">
|
||||
<Button
|
||||
type="text"
|
||||
icon={<EditOutlined />}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
onEdit(group);
|
||||
}}
|
||||
/>
|
||||
</Tooltip>
|
||||
<Tooltip title="Delete">
|
||||
<Button
|
||||
type="text"
|
||||
danger
|
||||
icon={<DeleteOutlined />}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
onDelete(group);
|
||||
}}
|
||||
/>
|
||||
</Tooltip>
|
||||
</Flex>
|
||||
),
|
||||
},
|
||||
];
|
||||
const toggleUsage = useCallback((group: RoutingGroup) => {
|
||||
setExpanded((previous) => {
|
||||
const current = previous === true ? {} : previous;
|
||||
return { ...current, [group.group_name]: current[group.group_name] !== true };
|
||||
});
|
||||
}, []);
|
||||
|
||||
const columns = useMemo(() => {
|
||||
const deps = { onEdit, onDelete, onToggleUsage: toggleUsage };
|
||||
return getRoutingGroupsTableColumns(deps);
|
||||
}, [onEdit, onDelete, toggleUsage]);
|
||||
|
||||
return (
|
||||
<Table<RoutingGroup>
|
||||
rowKey="group_name"
|
||||
<DataTable
|
||||
data={groups}
|
||||
columns={columns}
|
||||
dataSource={groups}
|
||||
loading={loading}
|
||||
pagination={false}
|
||||
expandable={{
|
||||
expandedRowKeys,
|
||||
onExpandedRowsChange: (keys) => setExpandedRowKeys([...keys]),
|
||||
expandedRowRender: (group) => (
|
||||
<div className="bg-gray-50 border border-gray-200 rounded-md p-4 my-2">
|
||||
<Flex align="center" gap={8} className="mb-2">
|
||||
<CodeOutlined className="text-blue-500" />
|
||||
<Text strong>How routing works for this group</Text>
|
||||
</Flex>
|
||||
<Paragraph className="text-sm text-gray-600 mb-3">
|
||||
Callers request any model in the group by name — LiteLLM picks a deployment behind the scenes using the{" "}
|
||||
<Text strong>{formatStrategyLabel(group.routing_strategy)}</Text> strategy.
|
||||
</Paragraph>
|
||||
<RoutingGroupSnippet group={group} baseUrl={baseUrl} />
|
||||
</div>
|
||||
),
|
||||
}}
|
||||
getRowId={(group) => group.group_name}
|
||||
sortingMode="client"
|
||||
sorting={sorting}
|
||||
onSortingChange={setSorting}
|
||||
expanded={expanded}
|
||||
onExpandedChange={setExpanded}
|
||||
getRowCanExpand={() => true}
|
||||
renderSubComponent={({ row }) => <RoutingGroupUsagePanel group={row.original} baseUrl={baseUrl} />}
|
||||
isLoading={isLoading}
|
||||
loadingMessage="Loading routing groups…"
|
||||
noDataMessage={<EmptyState />}
|
||||
size="compact"
|
||||
/>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -0,0 +1,111 @@
|
|||
"use client";
|
||||
|
||||
import type { ColumnDef } from "@tanstack/react-table";
|
||||
import { GitBranch, MoreHorizontal, Pencil, Trash2 } from "lucide-react";
|
||||
|
||||
import { DataTableSortHeader } from "@/components/shared/DataTable";
|
||||
import { IdentityCell, ModelsCell } from "@/components/shared/table_cells";
|
||||
import { buttonVariants } from "@/components/ui/button";
|
||||
import {
|
||||
DropdownMenu,
|
||||
DropdownMenuContent,
|
||||
DropdownMenuItem,
|
||||
DropdownMenuTrigger,
|
||||
} from "@/components/ui/dropdown-menu";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
|
||||
import { formatStrategyLabel } from "./strategy";
|
||||
import type { RoutingGroup } from "./types";
|
||||
|
||||
interface RoutingGroupRowActionsProps {
|
||||
group: RoutingGroup;
|
||||
onEdit: (group: RoutingGroup) => void;
|
||||
onDelete: (group: RoutingGroup) => void;
|
||||
}
|
||||
|
||||
function RoutingGroupRowActions({ group, onEdit, onDelete }: RoutingGroupRowActionsProps) {
|
||||
return (
|
||||
<DropdownMenu>
|
||||
<DropdownMenuTrigger
|
||||
aria-label={`Open actions for ${group.group_name}`}
|
||||
data-testid={`routing-group-actions-${group.group_name}`}
|
||||
className={cn(buttonVariants({ variant: "ghost", size: "icon-sm" }), "text-muted-foreground")}
|
||||
>
|
||||
<MoreHorizontal className="size-4" />
|
||||
</DropdownMenuTrigger>
|
||||
<DropdownMenuContent align="end" className="w-44">
|
||||
<DropdownMenuItem data-testid="routing-group-action-edit" onClick={() => onEdit(group)}>
|
||||
<Pencil />
|
||||
Edit
|
||||
</DropdownMenuItem>
|
||||
<DropdownMenuItem
|
||||
variant="destructive"
|
||||
data-testid="routing-group-action-delete"
|
||||
onClick={() => onDelete(group)}
|
||||
>
|
||||
<Trash2 />
|
||||
Delete
|
||||
</DropdownMenuItem>
|
||||
</DropdownMenuContent>
|
||||
</DropdownMenu>
|
||||
);
|
||||
}
|
||||
|
||||
interface RoutingGroupsTableColumnsDeps {
|
||||
onEdit: (group: RoutingGroup) => void;
|
||||
onDelete: (group: RoutingGroup) => void;
|
||||
onToggleUsage: (group: RoutingGroup) => void;
|
||||
}
|
||||
|
||||
export const getRoutingGroupsTableColumns = ({
|
||||
onEdit,
|
||||
onDelete,
|
||||
onToggleUsage,
|
||||
}: RoutingGroupsTableColumnsDeps): ColumnDef<RoutingGroup>[] => [
|
||||
{
|
||||
id: "group_name",
|
||||
accessorKey: "group_name",
|
||||
meta: { title: "Group Name", skeleton: "text" },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Group Name" />,
|
||||
size: 240,
|
||||
enableSorting: true,
|
||||
cell: ({ row }) => (
|
||||
<IdentityCell title={row.original.group_name} className="max-w-60" onClick={() => onToggleUsage(row.original)} />
|
||||
),
|
||||
},
|
||||
{
|
||||
id: "models",
|
||||
meta: { title: "Models", skeleton: "chips" },
|
||||
header: "Models",
|
||||
size: 320,
|
||||
enableSorting: false,
|
||||
cell: ({ row }) => <ModelsCell models={row.original.models} />,
|
||||
},
|
||||
{
|
||||
id: "routing_strategy",
|
||||
accessorKey: "routing_strategy",
|
||||
meta: { title: "Strategy", skeleton: "text" },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Strategy" />,
|
||||
size: 180,
|
||||
enableSorting: true,
|
||||
cell: ({ row }) => (
|
||||
<span className="flex items-center gap-1.5 text-sm">
|
||||
<GitBranch className="size-4 shrink-0 text-muted-foreground" />
|
||||
{formatStrategyLabel(row.original.routing_strategy)}
|
||||
</span>
|
||||
),
|
||||
},
|
||||
{
|
||||
id: "actions",
|
||||
meta: { className: "text-right", headerClassName: "text-right" },
|
||||
header: () => <span className="sr-only">Actions</span>,
|
||||
size: 64,
|
||||
enableSorting: false,
|
||||
enableHiding: false,
|
||||
cell: ({ row }) => (
|
||||
<div className="flex justify-end">
|
||||
<RoutingGroupRowActions group={row.original} onEdit={onEdit} onDelete={onDelete} />
|
||||
</div>
|
||||
),
|
||||
},
|
||||
];
|
||||
|
|
@ -126,7 +126,7 @@ const RoutingGroups: React.FC = () => {
|
|||
|
||||
<RoutingGroupsTable
|
||||
groups={filteredGroups}
|
||||
loading={isLoading}
|
||||
isLoading={isLoading}
|
||||
onEdit={openEdit}
|
||||
onDelete={(g) => setDeletingGroup(g)}
|
||||
proxyBaseUrl={proxySettings.LITELLM_UI_API_DOC_BASE_URL?.trim() || proxySettings.PROXY_BASE_URL || ""}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,8 @@
|
|||
const STRATEGY_LABELS: Readonly<Record<string, string>> = {
|
||||
"simple-shuffle": "Simple Shuffle",
|
||||
"least-busy": "Least Busy",
|
||||
"usage-based-routing": "Usage Based",
|
||||
"latency-based-routing": "Latency Based",
|
||||
};
|
||||
|
||||
export const formatStrategyLabel = (strategy: string): string => STRATEGY_LABELS[strategy] ?? strategy;
|
||||
|
|
@ -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