mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_guardrail_event_hook_resync
This commit is contained in:
commit
31a3c55737
31 changed files with 2504 additions and 46 deletions
|
|
@ -117,7 +117,7 @@
|
|||
"limit": 110
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 695
|
||||
"limit": 692
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 5
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import json
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, cast
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
|
|
@ -62,6 +63,31 @@ if TYPE_CHECKING:
|
|||
# --- typed sub-structures ---------------------------------------------------- #
|
||||
|
||||
|
||||
def _cache_token_value(*values: object) -> int | None:
|
||||
explicit_zero = False
|
||||
invalid_before_zero = False
|
||||
for raw_value in values:
|
||||
if raw_value is None:
|
||||
continue
|
||||
if isinstance(raw_value, bool):
|
||||
parsed = None
|
||||
else:
|
||||
try:
|
||||
parsed = as_int(raw_value)
|
||||
except (OverflowError, ValueError):
|
||||
parsed = None
|
||||
if parsed is None:
|
||||
if not explicit_zero:
|
||||
invalid_before_zero = True
|
||||
elif parsed > 0:
|
||||
return parsed
|
||||
elif parsed == 0:
|
||||
explicit_zero = True
|
||||
elif not explicit_zero:
|
||||
invalid_before_zero = True
|
||||
return 0 if explicit_zero and not invalid_before_zero else None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LLMRequestParams:
|
||||
temperature: float | None = None
|
||||
|
|
@ -104,12 +130,25 @@ class LLMUsage:
|
|||
metadata: Final[Mapping[str, object]] = payload.get("metadata") or {}
|
||||
raw_usage: Final = metadata.get("usage_object")
|
||||
usage_object: Final[Mapping[str, object]] = raw_usage if isinstance(raw_usage, Mapping) else {}
|
||||
raw_details: Final = usage_object.get("prompt_tokens_details")
|
||||
prompt_details: Final[Mapping[str, object]] = (
|
||||
raw_details if isinstance(raw_details, Mapping) else MappingProxyType({})
|
||||
)
|
||||
return cls(
|
||||
input_tokens=as_int(payload.get("prompt_tokens")),
|
||||
output_tokens=as_int(payload.get("completion_tokens")),
|
||||
total_tokens=as_int(payload.get("total_tokens")),
|
||||
cache_creation_input_tokens=as_int(usage_object.get("cache_creation_input_tokens")),
|
||||
cache_read_input_tokens=as_int(usage_object.get("cache_read_input_tokens")),
|
||||
cache_creation_input_tokens=_cache_token_value(
|
||||
usage_object.get("cache_creation_input_tokens"),
|
||||
prompt_details.get("cache_write_tokens"),
|
||||
prompt_details.get("cache_creation_tokens"),
|
||||
prompt_details.get("cache_creation_input_tokens"),
|
||||
),
|
||||
cache_read_input_tokens=_cache_token_value(
|
||||
usage_object.get("cache_read_input_tokens"),
|
||||
prompt_details.get("cached_tokens"),
|
||||
usage_object.get("prompt_cache_hit_tokens"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import math
|
|||
import os
|
||||
import sys
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import replace
|
||||
from datetime import datetime, timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast
|
||||
|
||||
|
|
@ -58,6 +59,7 @@ from litellm.types.utils import (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from prometheus_client import Gauge
|
||||
from prometheus_client.metrics import MetricWrapperBase
|
||||
|
||||
from litellm.router import Router
|
||||
|
|
@ -476,6 +478,30 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=self.get_labels_for_metric("litellm_remaining_api_key_tokens_for_model"),
|
||||
)
|
||||
|
||||
self.litellm_api_key_rate_limit_allowed_metric = self._gauge_factory(
|
||||
"litellm_api_key_rate_limit_allowed_metric",
|
||||
"Configured rate limit for the API Key in the current window (rpm_limit / tpm_limit), by rate_limit_type",
|
||||
labelnames=self.get_labels_for_metric("litellm_api_key_rate_limit_allowed_metric"),
|
||||
)
|
||||
|
||||
self.litellm_api_key_rate_limit_used_metric = self._gauge_factory(
|
||||
"litellm_api_key_rate_limit_used_metric",
|
||||
"Requests or tokens the API Key has consumed in the current rate limit window, by rate_limit_type",
|
||||
labelnames=self.get_labels_for_metric("litellm_api_key_rate_limit_used_metric"),
|
||||
)
|
||||
|
||||
self.litellm_team_rate_limit_allowed_metric = self._gauge_factory(
|
||||
"litellm_team_rate_limit_allowed_metric",
|
||||
"Configured rate limit for the Team in the current window (team rpm_limit / tpm_limit), by rate_limit_type",
|
||||
labelnames=self.get_labels_for_metric("litellm_team_rate_limit_allowed_metric"),
|
||||
)
|
||||
|
||||
self.litellm_team_rate_limit_used_metric = self._gauge_factory(
|
||||
"litellm_team_rate_limit_used_metric",
|
||||
"Requests or tokens the Team has consumed in the current rate limit window, by rate_limit_type",
|
||||
labelnames=self.get_labels_for_metric("litellm_team_rate_limit_used_metric"),
|
||||
)
|
||||
|
||||
########################################
|
||||
# LLM API Deployment Metrics / analytics
|
||||
########################################
|
||||
|
|
@ -1475,6 +1501,11 @@ class PrometheusLogger(CustomLogger):
|
|||
model_id=enum_values.model_id,
|
||||
)
|
||||
|
||||
self._set_key_and_team_rate_limit_metrics(
|
||||
standard_logging_payload=standard_logging_payload, # pyright: ignore[reportArgumentType] # isinstance(dict) above narrows the TypedDict to dict[Unknown, Unknown]
|
||||
enum_values=enum_values,
|
||||
)
|
||||
|
||||
# set latency metrics
|
||||
self._set_latency_metrics(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -2002,17 +2033,102 @@ class PrometheusLogger(CustomLogger):
|
|||
"""
|
||||
if standard_logging_payload is None:
|
||||
return None
|
||||
return PrometheusLogger._get_int_from_v3_rate_limit_headers(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
header_name=f"x-ratelimit-model_per_key-remaining-{rate_limit_type}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_int_from_v3_rate_limit_headers(
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
header_name: str,
|
||||
) -> int | None:
|
||||
hidden_params: Final = standard_logging_payload.get("hidden_params")
|
||||
if hidden_params is None:
|
||||
return None
|
||||
additional_headers: Final = hidden_params.get("additional_headers")
|
||||
additional_headers: Final[Mapping[str, object] | None] = hidden_params.get("additional_headers")
|
||||
if additional_headers is None:
|
||||
return None
|
||||
value: Final = dict(additional_headers).get(f"x-ratelimit-model_per_key-remaining-{rate_limit_type}")
|
||||
value: Final = additional_headers.get(header_name)
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
return None
|
||||
return value
|
||||
|
||||
def _set_key_and_team_rate_limit_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
) -> None:
|
||||
"""
|
||||
Export the key-level and team-level RPM / TPM limit and current window
|
||||
usage from the ``x-ratelimit-{api_key,team}-{limit,remaining}-*``
|
||||
headers the v3 rate limiter mirrors into the logging payload. The
|
||||
limiter already read these counters (from Redis when configured) on
|
||||
the request path, so no extra store lookup happens here. Descriptors
|
||||
without a configured limit emit no header, so their series is removed
|
||||
rather than left at the value from before the limit was dropped.
|
||||
"""
|
||||
descriptor_gauges: Final[
|
||||
tuple[tuple[Literal["api_key", "team"], DEFINED_PROMETHEUS_METRICS, Gauge, Gauge], ...]
|
||||
] = (
|
||||
(
|
||||
"api_key",
|
||||
"litellm_api_key_rate_limit_allowed_metric",
|
||||
self.litellm_api_key_rate_limit_allowed_metric,
|
||||
self.litellm_api_key_rate_limit_used_metric,
|
||||
),
|
||||
(
|
||||
"team",
|
||||
"litellm_team_rate_limit_allowed_metric",
|
||||
self.litellm_team_rate_limit_allowed_metric,
|
||||
self.litellm_team_rate_limit_used_metric,
|
||||
),
|
||||
)
|
||||
for descriptor_key, metric_name, allowed_gauge, used_gauge in descriptor_gauges:
|
||||
for rate_limit_type in ("requests", "tokens"):
|
||||
self._set_rate_limit_allowed_and_used_gauges(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
descriptor_key=descriptor_key,
|
||||
metric_name=metric_name,
|
||||
allowed_gauge=allowed_gauge,
|
||||
used_gauge=used_gauge,
|
||||
rate_limit_type=rate_limit_type,
|
||||
)
|
||||
|
||||
def _set_rate_limit_allowed_and_used_gauges(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
descriptor_key: Literal["api_key", "team"],
|
||||
metric_name: DEFINED_PROMETHEUS_METRICS,
|
||||
allowed_gauge: Gauge,
|
||||
used_gauge: Gauge,
|
||||
rate_limit_type: Literal["requests", "tokens"],
|
||||
) -> None:
|
||||
limit: Final = self._get_int_from_v3_rate_limit_headers(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
header_name=f"x-ratelimit-{descriptor_key}-limit-{rate_limit_type}",
|
||||
)
|
||||
remaining: Final = self._get_int_from_v3_rate_limit_headers(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
header_name=f"x-ratelimit-{descriptor_key}-remaining-{rate_limit_type}",
|
||||
)
|
||||
labelled_values: Final = replace(enum_values, rate_limit_type=rate_limit_type)
|
||||
labelnames: Final = self.get_labels_for_metric(metric_name)
|
||||
labels: Final = prometheus_label_factory(
|
||||
supported_enum_labels=labelnames,
|
||||
enum_values=labelled_values,
|
||||
label_context=PrometheusLabelFactoryContext(labelled_values),
|
||||
)
|
||||
if limit is None or remaining is None:
|
||||
label_values: Final = tuple(labels.get(label) for label in labelnames)
|
||||
self._bounded_prometheus_series_tracker.remove_series(allowed_gauge, label_values)
|
||||
self._bounded_prometheus_series_tracker.remove_series(used_gauge, label_values)
|
||||
return
|
||||
allowed_gauge.labels(**labels).set(limit)
|
||||
used_gauge.labels(**labels).set(limit - remaining)
|
||||
|
||||
def _set_virtual_key_rate_limit_metrics(
|
||||
self,
|
||||
user_api_key: str | None,
|
||||
|
|
|
|||
|
|
@ -60,6 +60,10 @@ class BoundedPrometheusSeriesTracker:
|
|||
break
|
||||
del series[tracked_label_values]
|
||||
|
||||
def remove_series(self, metric: object, label_values: tuple[str | None, ...]) -> bool:
|
||||
"""Drop one child series, True when it is gone (removed or never existed)."""
|
||||
return self._remove_metric_child(metric, label_values)
|
||||
|
||||
def _should_run_ttl_cleanup(
|
||||
self,
|
||||
metric_name: str,
|
||||
|
|
|
|||
|
|
@ -155,8 +155,8 @@ class BaseTranslation(ABC):
|
|||
self,
|
||||
exc: "ModifyResponseException",
|
||||
stream_started: bool = False,
|
||||
responses_so_far: list[Any] | None = None,
|
||||
) -> list[bytes] | None:
|
||||
responses_so_far: Sequence[Any] | None = None,
|
||||
) -> Sequence[bytes] | None:
|
||||
"""
|
||||
Build the streaming chunks that deliver a guardrail block message and
|
||||
cleanly terminate the stream in this provider's wire format.
|
||||
|
|
|
|||
|
|
@ -124,6 +124,61 @@ def blocked_responses_api_usage(original_response: object) -> ResponseAPIUsage:
|
|||
)
|
||||
|
||||
|
||||
def stream_item_field(item: object, field: str) -> object | None:
|
||||
if isinstance(item, dict):
|
||||
return item.get(field)
|
||||
return getattr(item, field, None)
|
||||
|
||||
|
||||
def blocked_chat_stream_usage(original_response: object) -> tuple[int, int]:
|
||||
"""
|
||||
``(prompt_tokens, completion_tokens)`` for a synthetic guardrail-blocked
|
||||
chat completions stream.
|
||||
|
||||
A mid-stream block carries the chunks received so far as a list; real usage
|
||||
rides on the final chunk when the upstream sent one
|
||||
(``stream_options.include_usage``). Non-list originals defer to
|
||||
``blocked_response_usage``.
|
||||
"""
|
||||
if not isinstance(original_response, list):
|
||||
usage: Final = blocked_response_usage(original_response)
|
||||
return usage.get("input_tokens", 0), usage.get("output_tokens", 0)
|
||||
usage_obj: Final = next(
|
||||
(
|
||||
chunk_usage
|
||||
for item in reversed(original_response)
|
||||
if (chunk_usage := stream_item_field(item, "usage")) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
return (
|
||||
_usage_tokens(usage_obj, "prompt_tokens", "input_tokens"),
|
||||
_usage_tokens(usage_obj, "completion_tokens", "output_tokens"),
|
||||
)
|
||||
|
||||
|
||||
def blocked_responses_stream_usage(original_response: object) -> ResponseAPIUsage:
|
||||
"""
|
||||
``ResponseAPIUsage`` for a synthetic guardrail-blocked /v1/responses stream.
|
||||
|
||||
A mid-stream block carries the events received so far as a list; real usage
|
||||
rides on the ``response.completed`` event's response when the upstream sent
|
||||
one. Non-list originals defer to ``blocked_responses_api_usage``.
|
||||
"""
|
||||
if not isinstance(original_response, list):
|
||||
return blocked_responses_api_usage(original_response)
|
||||
completed: Final = next(
|
||||
(
|
||||
response
|
||||
for item in reversed(original_response)
|
||||
if stream_item_field(item, "type") == "response.completed"
|
||||
and (response := stream_item_field(item, "response")) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
return blocked_responses_api_usage(completed)
|
||||
|
||||
|
||||
def effective_skip_system_message_for_guardrail(guardrail_to_apply: Any) -> bool:
|
||||
per: Final = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None)
|
||||
if per is not None:
|
||||
|
|
|
|||
|
|
@ -14,9 +14,14 @@ Pattern Overview:
|
|||
This pattern can be replicated for other message formats (e.g., Anthropic).
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
||||
|
|
@ -24,6 +29,7 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
|||
StreamTransformSink,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_chat_stream_usage,
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
|
|
@ -32,6 +38,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
openai_tool_name,
|
||||
role_out_of_guardrail_scope,
|
||||
scoped_structured_message_indices,
|
||||
stream_item_field,
|
||||
)
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
|
|
@ -49,7 +56,10 @@ from litellm.types.utils import (
|
|||
if TYPE_CHECKING:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
|
@ -1005,3 +1015,129 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
else:
|
||||
# Subsequent chunks - clear the text
|
||||
content_item["text"] = ""
|
||||
|
||||
def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool:
|
||||
"""
|
||||
True once any relayed chunk carries a non-null ``finish_reason``.
|
||||
|
||||
The unified guardrail's ``end_of_stream_only`` streaming path probes
|
||||
this via ``hasattr`` to withhold the terminal chunks until
|
||||
end-of-stream moderation runs, so a block can replace the finish
|
||||
instead of trailing after a ``finish_reason`` the client already saw.
|
||||
"""
|
||||
return any(
|
||||
stream_item_field(choice, "finish_reason") is not None
|
||||
for item in responses_so_far
|
||||
for choice in _stream_chunk_choices(item)
|
||||
)
|
||||
|
||||
def build_block_sse_chunks(
|
||||
self,
|
||||
exc: "ModifyResponseException",
|
||||
stream_started: bool = False,
|
||||
responses_so_far: Sequence[object] | None = None,
|
||||
) -> Sequence[bytes]:
|
||||
"""
|
||||
Build OpenAI chat-completions SSE chunks that deliver the guardrail
|
||||
block message and terminate the stream cleanly, mirroring the
|
||||
non-streaming block response: ``finish_reason`` ``content_filter`` plus
|
||||
the real usage the upstream call consumed.
|
||||
|
||||
- ``stream_started`` False (buffered / pre-stream): nothing has been
|
||||
sent, so open a standalone completion with a ``role`` delta.
|
||||
- ``stream_started`` True (sampling / mid-stream): chunks already
|
||||
reached the client, so continue the in-progress completion (reuse its
|
||||
id/created/model, content-only delta).
|
||||
|
||||
The proxy's data generator appends ``data: [DONE]`` itself.
|
||||
"""
|
||||
chunk_id, created, model = _blocked_stream_identity(exc, responses_so_far or ())
|
||||
prompt_tokens, completion_tokens = blocked_chat_stream_usage(exc.original_response)
|
||||
continuation_delta: Final[_BlockedChunkDelta] = {"content": exc.message}
|
||||
standalone_delta: Final[_BlockedChunkDelta] = {"role": "assistant", "content": exc.message}
|
||||
message_chunk: Final[_BlockedChunk] = {
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": (
|
||||
{
|
||||
"index": 0,
|
||||
"delta": continuation_delta if stream_started else standalone_delta,
|
||||
"finish_reason": None,
|
||||
},
|
||||
),
|
||||
}
|
||||
final_chunk: Final[_BlockedChunk] = {
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": ({"index": 0, "delta": {}, "finish_reason": "content_filter"},),
|
||||
"usage": {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
},
|
||||
}
|
||||
return _chat_sse_chunk(message_chunk), _chat_sse_chunk(final_chunk)
|
||||
|
||||
|
||||
class _BlockedChunkDelta(TypedDict, total=False):
|
||||
role: ReadOnly[str]
|
||||
content: ReadOnly[str]
|
||||
|
||||
|
||||
class _BlockedChunkChoice(TypedDict):
|
||||
index: ReadOnly[int]
|
||||
delta: ReadOnly[_BlockedChunkDelta]
|
||||
finish_reason: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _BlockedChunkUsage(TypedDict):
|
||||
prompt_tokens: ReadOnly[int]
|
||||
completion_tokens: ReadOnly[int]
|
||||
total_tokens: ReadOnly[int]
|
||||
|
||||
|
||||
class _BlockedChunk(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
object: ReadOnly[str]
|
||||
created: ReadOnly[int]
|
||||
model: ReadOnly[str]
|
||||
choices: ReadOnly[tuple[_BlockedChunkChoice, ...]]
|
||||
usage: NotRequired[ReadOnly[_BlockedChunkUsage]]
|
||||
|
||||
|
||||
def _chat_sse_chunk(payload: _BlockedChunk) -> bytes:
|
||||
return f"data: {json.dumps(payload)}\n\n".encode()
|
||||
|
||||
|
||||
def _stream_chunk_choices(item: object) -> Sequence[object]:
|
||||
choices: Final = stream_item_field(item, "choices")
|
||||
if isinstance(choices, Sequence) and not isinstance(choices, (str, bytes)):
|
||||
return choices
|
||||
return ()
|
||||
|
||||
|
||||
def _blocked_stream_identity(
|
||||
exc: "ModifyResponseException", responses_so_far: Sequence[object]
|
||||
) -> tuple[str, int, str]:
|
||||
identified: Final = next(
|
||||
(
|
||||
(chunk_id, item)
|
||||
for item in responses_so_far
|
||||
if isinstance(chunk_id := stream_item_field(item, "id"), str) and chunk_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if identified is None:
|
||||
return f"chatcmpl-{uuid.uuid4()}", int(time.time()), exc.model
|
||||
chunk_id, source = identified
|
||||
created: Final = stream_item_field(source, "created")
|
||||
model: Final = stream_item_field(source, "model")
|
||||
return (
|
||||
chunk_id,
|
||||
created if isinstance(created, int) else int(time.time()),
|
||||
model if isinstance(model, str) and model else exc.model,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -28,12 +28,16 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
|
|||
- text: str
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||
|
||||
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
|
||||
from openai.types.responses.tool_param import FunctionToolParam
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -41,17 +45,33 @@ from litellm.completion_extras.litellm_responses_transformation.transformation i
|
|||
OpenAiResponsesToChatCompletionStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_responses_stream_usage,
|
||||
stream_item_field,
|
||||
)
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
BaseLiteLLMOpenAIResponseObject,
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
ContentPartAddedEvent,
|
||||
ContentPartDoneEvent,
|
||||
ContentPartDonePartOutputText,
|
||||
ErrorEvent,
|
||||
ErrorEventError,
|
||||
OpenAIMcpServerTool,
|
||||
OutputItemAddedEvent,
|
||||
OutputItemDoneEvent,
|
||||
OutputTextDeltaEvent,
|
||||
OutputTextDoneEvent,
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
|
|
@ -63,11 +83,13 @@ from litellm.types.utils import GenericGuardrailAPIInputs
|
|||
if TYPE_CHECKING:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import ResponseInputParam
|
||||
from litellm.types.utils import ResponsesAPIResponse
|
||||
|
||||
|
||||
class ResponseOutputEnvelope(TypedDict, total=False):
|
||||
|
|
@ -865,3 +887,331 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
content[content_idx]["text"] = guardrail_response
|
||||
elif hasattr(content[content_idx], "text"):
|
||||
content[content_idx].text = guardrail_response
|
||||
|
||||
def build_block_sse_chunks(
|
||||
self,
|
||||
exc: "ModifyResponseException",
|
||||
stream_started: bool = False,
|
||||
responses_so_far: Sequence[object] | None = None,
|
||||
) -> Sequence[bytes]:
|
||||
"""
|
||||
Build Responses API SSE events that deliver the guardrail block message
|
||||
and terminate the stream cleanly, mirroring the non-streaming block
|
||||
response: a completed response whose only output is the violation text,
|
||||
with the real usage the upstream call consumed.
|
||||
|
||||
- ``stream_started`` False (buffered / pre-stream): nothing has been
|
||||
sent, so emit the full synthetic sequence (``response.created``
|
||||
through ``response.completed``).
|
||||
- ``stream_started`` True (sampling / mid-stream): events already
|
||||
reached the client, so continue the in-progress response: close the
|
||||
output item still open on the wire, deliver the block message as a
|
||||
new output item under the same response id, and close with a
|
||||
``response.completed`` carrying only the replacement item.
|
||||
|
||||
The proxy's data generator appends ``data: [DONE]`` itself.
|
||||
"""
|
||||
events: Final = (
|
||||
self._block_continuation_events(exc, responses_so_far or ())
|
||||
if stream_started
|
||||
else self._standalone_block_events(exc)
|
||||
)
|
||||
return tuple(
|
||||
f"data: {event.model_dump_json(exclude_none=True, exclude_unset=True, serialize_as_any=True)}\n\n".encode()
|
||||
for event in events
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _standalone_block_events(exc: "ModifyResponseException") -> Sequence[ResponsesAPIStreamingResponse]:
|
||||
from litellm.responses.streaming_iterator import build_synthetic_response_events
|
||||
|
||||
return build_synthetic_response_events(
|
||||
transformed=_blocked_response(exc, response_id=f"resp_{uuid.uuid4()}", model=exc.model),
|
||||
logging_obj=None,
|
||||
chunk_size=max(len(exc.message), 1),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _block_continuation_events(
|
||||
exc: "ModifyResponseException", responses_so_far: Sequence[object]
|
||||
) -> Sequence[ResponsesAPIStreamingResponse]:
|
||||
response_id, model, output_index = _continuation_identity(exc, responses_so_far)
|
||||
item: Final = _blocked_output_item(exc)
|
||||
item_id: Final = item.id
|
||||
part: Final[_BlockedContentPart] = {"type": "output_text", "text": exc.message, "annotations": ()}
|
||||
done_part: Final[_BlockedDoneContentPart] = {
|
||||
"type": "output_text",
|
||||
"text": exc.message,
|
||||
"annotations": (),
|
||||
"logprobs": None,
|
||||
}
|
||||
return (
|
||||
*_open_item_closing_events(responses_so_far),
|
||||
OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=output_index,
|
||||
item=item,
|
||||
),
|
||||
ContentPartAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
|
||||
item_id=item_id,
|
||||
output_index=output_index,
|
||||
content_index=0,
|
||||
part=BaseLiteLLMOpenAIResponseObject.model_validate(part),
|
||||
),
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id=item_id,
|
||||
output_index=output_index,
|
||||
content_index=0,
|
||||
delta=exc.message,
|
||||
),
|
||||
OutputTextDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
|
||||
item_id=item_id,
|
||||
output_index=output_index,
|
||||
content_index=0,
|
||||
text=exc.message,
|
||||
),
|
||||
ContentPartDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_DONE,
|
||||
item_id=item_id,
|
||||
output_index=output_index,
|
||||
content_index=0,
|
||||
part=ContentPartDonePartOutputText.model_validate(done_part),
|
||||
),
|
||||
OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=output_index,
|
||||
item=item,
|
||||
),
|
||||
ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=_blocked_response(exc, response_id=response_id, model=model, output_item=item),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class _BlockedContentPart(TypedDict):
|
||||
type: ReadOnly[str]
|
||||
text: ReadOnly[str]
|
||||
annotations: ReadOnly[tuple[object, ...]]
|
||||
|
||||
|
||||
class _BlockedDoneContentPart(TypedDict):
|
||||
type: ReadOnly[str]
|
||||
text: ReadOnly[str]
|
||||
annotations: ReadOnly[tuple[object, ...]]
|
||||
logprobs: ReadOnly[None]
|
||||
|
||||
|
||||
class _BlockedItemPayload(TypedDict):
|
||||
type: ReadOnly[str]
|
||||
id: ReadOnly[str]
|
||||
status: ReadOnly[str]
|
||||
role: ReadOnly[str]
|
||||
content: ReadOnly[tuple[_BlockedContentPart, ...]]
|
||||
|
||||
|
||||
class _BlockedResponsePayload(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
object: ReadOnly[str]
|
||||
created_at: ReadOnly[int]
|
||||
model: ReadOnly[str]
|
||||
output: ReadOnly[tuple[GenericResponseOutputItem, ...]]
|
||||
status: ReadOnly[str]
|
||||
usage: ReadOnly[ResponseAPIUsage]
|
||||
|
||||
|
||||
def _blocked_output_item(exc: "ModifyResponseException") -> GenericResponseOutputItem:
|
||||
payload: Final[_BlockedItemPayload] = {
|
||||
"type": "message",
|
||||
"id": f"msg_{uuid.uuid4()}",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": ({"type": "output_text", "text": exc.message, "annotations": ()},),
|
||||
}
|
||||
return GenericResponseOutputItem.model_validate(payload)
|
||||
|
||||
|
||||
def _blocked_response(
|
||||
exc: "ModifyResponseException",
|
||||
response_id: str,
|
||||
model: str,
|
||||
output_item: GenericResponseOutputItem | None = None,
|
||||
) -> ResponsesAPIResponse:
|
||||
payload: Final[_BlockedResponsePayload] = {
|
||||
"id": response_id,
|
||||
"object": "response",
|
||||
"created_at": int(time.time()),
|
||||
"model": model,
|
||||
"output": (output_item if output_item is not None else _blocked_output_item(exc),),
|
||||
"status": "completed",
|
||||
"usage": blocked_responses_stream_usage(exc.original_response),
|
||||
}
|
||||
return ResponsesAPIResponse.model_validate(payload)
|
||||
|
||||
|
||||
def _continuation_identity(exc: "ModifyResponseException", responses_so_far: Sequence[object]) -> tuple[str, str, int]:
|
||||
responses: Final = tuple(
|
||||
response for item in responses_so_far if (response := stream_item_field(item, "response")) is not None
|
||||
)
|
||||
response_id: Final = next(
|
||||
(rid for response in responses if isinstance(rid := stream_item_field(response, "id"), str) and rid),
|
||||
f"resp_{uuid.uuid4()}",
|
||||
)
|
||||
model: Final = next(
|
||||
(m for response in responses if isinstance(m := stream_item_field(response, "model"), str) and m),
|
||||
exc.model,
|
||||
)
|
||||
indices: Final = tuple(
|
||||
index for item in responses_so_far if isinstance(index := stream_item_field(item, "output_index"), int)
|
||||
)
|
||||
return response_id, model, max(indices) + 1 if indices else 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _OpenItemState:
|
||||
item_id: str
|
||||
item_type: str
|
||||
role: str
|
||||
output_index: int
|
||||
content_index: int
|
||||
text: str
|
||||
part_open: bool
|
||||
payload: object
|
||||
|
||||
|
||||
def _open_item_state(responses_so_far: Sequence[object]) -> _OpenItemState | None:
|
||||
typed: Final = tuple((stream_item_field(event, "type"), event) for event in responses_so_far)
|
||||
added: Final = tuple(
|
||||
(added_index, stream_item_field(event, "item"))
|
||||
for event_type, event in typed
|
||||
if event_type == "response.output_item.added"
|
||||
and isinstance(added_index := stream_item_field(event, "output_index"), int)
|
||||
)
|
||||
done_indices: Final = frozenset(
|
||||
done_index
|
||||
for event_type, event in typed
|
||||
if event_type == "response.output_item.done"
|
||||
and isinstance(done_index := stream_item_field(event, "output_index"), int)
|
||||
)
|
||||
open_added: Final = tuple((index, payload) for index, payload in added if index not in done_indices)
|
||||
if not open_added:
|
||||
return None
|
||||
output_index, item_payload = open_added[-1]
|
||||
if item_payload is None:
|
||||
return None
|
||||
item_id: Final = stream_item_field(item_payload, "id")
|
||||
if not isinstance(item_id, str) or not item_id:
|
||||
return None
|
||||
raw_type: Final = stream_item_field(item_payload, "type")
|
||||
raw_role: Final = stream_item_field(item_payload, "role")
|
||||
part_added: Final = tuple(
|
||||
part_index
|
||||
for event_type, event in typed
|
||||
if event_type == "response.content_part.added"
|
||||
and stream_item_field(event, "item_id") == item_id
|
||||
and isinstance(part_index := stream_item_field(event, "content_index"), int)
|
||||
)
|
||||
part_done: Final = frozenset(
|
||||
part_done_index
|
||||
for event_type, event in typed
|
||||
if event_type == "response.content_part.done"
|
||||
and stream_item_field(event, "item_id") == item_id
|
||||
and isinstance(part_done_index := stream_item_field(event, "content_index"), int)
|
||||
)
|
||||
open_parts: Final = tuple(index for index in part_added if index not in part_done)
|
||||
text: Final = "".join(
|
||||
delta
|
||||
for event_type, event in typed
|
||||
if event_type == "response.output_text.delta"
|
||||
and stream_item_field(event, "item_id") == item_id
|
||||
and isinstance(delta := stream_item_field(event, "delta"), str)
|
||||
)
|
||||
return _OpenItemState(
|
||||
item_id=item_id,
|
||||
item_type=raw_type if isinstance(raw_type, str) and raw_type else "message",
|
||||
role=raw_role if isinstance(raw_role, str) and raw_role else "assistant",
|
||||
output_index=output_index,
|
||||
content_index=open_parts[-1] if open_parts else 0,
|
||||
text=text,
|
||||
part_open=bool(open_parts),
|
||||
payload=item_payload,
|
||||
)
|
||||
|
||||
|
||||
_item_fields_adapter: Final = TypeAdapter(Mapping[str, object])
|
||||
_no_item_fields: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _incomplete_item_fields(payload: object) -> Mapping[str, object]:
|
||||
raw: Final = payload.model_dump() if isinstance(payload, BaseModel) else payload
|
||||
if not isinstance(raw, dict):
|
||||
return _no_item_fields
|
||||
return _item_fields_adapter.validate_python(raw)
|
||||
|
||||
|
||||
def _open_item_closing_events(responses_so_far: Sequence[object]) -> Sequence[ResponsesAPIStreamingResponse]:
|
||||
"""Close the output item still in progress on the relayed stream before the
|
||||
block item is appended: strict Responses clients reject a
|
||||
``response.completed`` that arrives while an earlier ``output_item.added``
|
||||
was never closed. A message item closes ``completed`` with exactly the text
|
||||
the client has received so far; any other item type (a function call the
|
||||
guardrail rejected, for instance) closes ``incomplete`` so the synthetic
|
||||
done event can never authorize acting on it."""
|
||||
open_item: Final = _open_item_state(responses_so_far)
|
||||
if open_item is None:
|
||||
return ()
|
||||
if open_item.item_type != "message":
|
||||
return (
|
||||
OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=open_item.output_index,
|
||||
item=BaseLiteLLMOpenAIResponseObject.model_validate(
|
||||
MappingProxyType({**_incomplete_item_fields(open_item.payload), "status": "incomplete"})
|
||||
),
|
||||
),
|
||||
)
|
||||
partial_part: Final[_BlockedContentPart] = {
|
||||
"type": "output_text",
|
||||
"text": open_item.text,
|
||||
"annotations": (),
|
||||
}
|
||||
closed_payload: Final[_BlockedItemPayload] = {
|
||||
"type": open_item.item_type,
|
||||
"id": open_item.item_id,
|
||||
"status": "completed",
|
||||
"role": open_item.role,
|
||||
"content": (partial_part,),
|
||||
}
|
||||
item_done: Final = OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=open_item.output_index,
|
||||
item=GenericResponseOutputItem.model_validate(closed_payload),
|
||||
)
|
||||
if not open_item.part_open:
|
||||
return (item_done,)
|
||||
partial_done_part: Final[_BlockedDoneContentPart] = {
|
||||
"type": "output_text",
|
||||
"text": open_item.text,
|
||||
"annotations": (),
|
||||
"logprobs": None,
|
||||
}
|
||||
return (
|
||||
OutputTextDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
|
||||
item_id=open_item.item_id,
|
||||
output_index=open_item.output_index,
|
||||
content_index=open_item.content_index,
|
||||
text=open_item.text,
|
||||
),
|
||||
ContentPartDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_DONE,
|
||||
item_id=open_item.item_id,
|
||||
output_index=open_item.output_index,
|
||||
content_index=open_item.content_index,
|
||||
part=ContentPartDonePartOutputText.model_validate(partial_done_part),
|
||||
),
|
||||
item_done,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ Canonical definition for ``litellm_usertable``. Re-exported from
|
|||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import ConfigDict, Field, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
from litellm.models.organization_membership import (
|
||||
|
|
@ -67,3 +67,11 @@ class LiteLLM_UserTable(LiteLLMPydanticObjectBase):
|
|||
if not self.models:
|
||||
return True
|
||||
return model_name in self.models
|
||||
|
||||
|
||||
class SCIMPlaceholder(BaseModel):
|
||||
"""A user row keyed by a value that names another account by SSO identity or email."""
|
||||
|
||||
placeholder_user_id: str
|
||||
resolved_user_ids: tuple[str, ...]
|
||||
team_ids: tuple[str, ...]
|
||||
|
|
|
|||
|
|
@ -32002,6 +32002,62 @@
|
|||
"title": "SCIMPatchOperation",
|
||||
"type": "object"
|
||||
},
|
||||
"SCIMPlaceholder": {
|
||||
"description": "A user row keyed by a value that names another account by SSO identity or email.",
|
||||
"properties": {
|
||||
"placeholder_user_id": {
|
||||
"title": "Placeholder User Id",
|
||||
"type": "string"
|
||||
},
|
||||
"resolved_user_ids": {
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"title": "Resolved User Ids",
|
||||
"type": "array"
|
||||
},
|
||||
"team_ids": {
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"title": "Team Ids",
|
||||
"type": "array"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"placeholder_user_id",
|
||||
"resolved_user_ids",
|
||||
"team_ids"
|
||||
],
|
||||
"title": "SCIMPlaceholder",
|
||||
"type": "object"
|
||||
},
|
||||
"SCIMPlaceholderMergeResult": {
|
||||
"properties": {
|
||||
"merged_into_user_id": {
|
||||
"title": "Merged Into User Id",
|
||||
"type": "string"
|
||||
},
|
||||
"placeholder_user_id": {
|
||||
"title": "Placeholder User Id",
|
||||
"type": "string"
|
||||
},
|
||||
"team_ids": {
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"title": "Team Ids",
|
||||
"type": "array"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"placeholder_user_id",
|
||||
"merged_into_user_id",
|
||||
"team_ids"
|
||||
],
|
||||
"title": "SCIMPlaceholderMergeResult",
|
||||
"type": "object"
|
||||
},
|
||||
"SCIMServiceProviderConfig": {
|
||||
"properties": {
|
||||
"authenticationSchemes": {
|
||||
|
|
@ -33641,6 +33697,129 @@
|
|||
"scim"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/scim/v2/placeholders": {
|
||||
"get": {
|
||||
"description": "List user rows whose id is another account's SSO identity or email.\n\nAn earlier release provisioned a group member it could not match as a user keyed\nby the raw member value, and that row now shadows the account the value really\nnames, so every push of that member is refused. This lists those rows so an\noperator can fold each one into the account it shadows with\n``POST /scim/v2/placeholders/{user_id}/merge``. A row that has an SSO identity of\nits own or owns virtual keys is left out: someone uses that account.",
|
||||
"operationId": "list_placeholders_scim_v2_placeholders_get",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "query",
|
||||
"name": "feature",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Feature"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/SCIMPlaceholder"
|
||||
},
|
||||
"title": "Response List Placeholders Scim V2 Placeholders Get",
|
||||
"type": "array"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "List Placeholders",
|
||||
"tags": [
|
||||
"scim"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/scim/v2/placeholders/{user_id}/merge": {
|
||||
"post": {
|
||||
"description": "Fold a placeholder user into the one account its id names by SSO identity or email.\n\nThe account is added to every team the placeholder is on, then the placeholder is\ndeleted the way ``DELETE /scim/v2/Users/{id}`` deletes a user, so the next group\npush resolves the member value to the real account. Refused with 409 when the row\nhas an SSO identity of its own, owns virtual keys, or names no account or several.",
|
||||
"operationId": "merge_placeholder_scim_v2_placeholders__user_id__merge_post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "user_id",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "User ID",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
{
|
||||
"in": "query",
|
||||
"name": "feature",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Feature"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/SCIMPlaceholderMergeResult"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Merge Placeholder",
|
||||
"tags": [
|
||||
"scim"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.models.user import SCIMPlaceholder
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_UserTable,
|
||||
|
|
@ -1862,6 +1863,89 @@ async def delete_user(
|
|||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
@scim_router.get(
|
||||
"/placeholders",
|
||||
response_model=tuple[SCIMPlaceholder, ...],
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
)
|
||||
async def list_placeholders() -> tuple[SCIMPlaceholder, ...]:
|
||||
"""
|
||||
List user rows whose id is another account's SSO identity or email.
|
||||
|
||||
An earlier release provisioned a group member it could not match as a user keyed
|
||||
by the raw member value, and that row now shadows the account the value really
|
||||
names, so every push of that member is refused. This lists those rows so an
|
||||
operator can fold each one into the account it shadows with
|
||||
``POST /scim/v2/placeholders/{user_id}/merge``. A row that has an SSO identity of
|
||||
its own or owns virtual keys is left out: someone uses that account.
|
||||
"""
|
||||
try:
|
||||
prisma_client: Final = await _get_prisma_client_or_raise_exception()
|
||||
async with prisma_client.tx() as tx:
|
||||
return await UserRepository(prisma_client).find_shadowing_placeholders(tx)
|
||||
except Exception as e:
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
def _placeholder_rejection(placeholder: LiteLLM_UserTable, resolved: tuple[str, ...], key_count: int) -> str | None:
|
||||
if placeholder.sso_user_id is not None:
|
||||
return f"User '{placeholder.user_id}' has an SSO identity of its own, so it is an account someone signs in to"
|
||||
if key_count:
|
||||
return f"User '{placeholder.user_id}' owns {key_count} virtual keys. Move or delete them before merging it"
|
||||
if not resolved:
|
||||
return f"User '{placeholder.user_id}' shadows no account: no other user has that id as SSO identity or email"
|
||||
if len(resolved) > 1:
|
||||
return (
|
||||
f"User '{placeholder.user_id}' names {len(resolved)} accounts ({', '.join(resolved)}). Resolve that first"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
@scim_router.post(
|
||||
"/placeholders/{user_id}/merge",
|
||||
response_model=SCIMPlaceholderMergeResult,
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
)
|
||||
async def merge_placeholder(
|
||||
user_id: str = Path(..., title="User ID"),
|
||||
) -> SCIMPlaceholderMergeResult:
|
||||
"""
|
||||
Fold a placeholder user into the one account its id names by SSO identity or email.
|
||||
|
||||
The account is added to every team the placeholder is on, then the placeholder is
|
||||
deleted the way ``DELETE /scim/v2/Users/{id}`` deletes a user, so the next group
|
||||
push resolves the member value to the real account. Refused with 409 when the row
|
||||
has an SSO identity of its own, owns virtual keys, or names no account or several.
|
||||
"""
|
||||
try:
|
||||
prisma_client: Final = await _get_prisma_client_or_raise_exception()
|
||||
placeholder: Final = await _check_user_exists(user_id)
|
||||
resolved: Final = tuple(
|
||||
other for other in await _users_named_by_member_value(user_id, prisma_client, take=None) if other != user_id
|
||||
)
|
||||
owned_keys: Final[_UserIdWhere] = {"user_id": user_id}
|
||||
keys: Final = await _table(VerificationTokenRepository(prisma_client)).find_many(where=owned_keys)
|
||||
rejection: Final = _placeholder_rejection(placeholder, resolved, len(keys))
|
||||
if rejection is not None:
|
||||
detail: Final[_ScimErrorDetail] = {"error": rejection}
|
||||
raise HTTPException(status_code=409, detail=detail)
|
||||
|
||||
target_user_id: Final = resolved[0]
|
||||
team_ids: Final = tuple(placeholder.teams)
|
||||
for team_id in team_ids:
|
||||
await _add_user_to_team(user_id=target_user_id, team_id=team_id)
|
||||
await delete_user(user_id=user_id)
|
||||
await _recompute_scim_member_roles(prisma_client, (target_user_id,))
|
||||
verbose_proxy_logger.info(
|
||||
"SCIM: merged placeholder user '%s' into '%s', moving teams %s", user_id, target_user_id, team_ids
|
||||
)
|
||||
return SCIMPlaceholderMergeResult(
|
||||
placeholder_user_id=user_id, merged_into_user_id=target_user_id, team_ids=team_ids
|
||||
)
|
||||
except Exception as e:
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
def _parse_member_entry(entry: object) -> SCIMMember | None:
|
||||
"""Parse one entry of a SCIM patch value, or None when it carries no id."""
|
||||
if isinstance(entry, str):
|
||||
|
|
|
|||
|
|
@ -812,6 +812,8 @@ def _resolve_team_callback_wiring(
|
|||
else { # mutable-ok: Logging arg
|
||||
**callback_vars,
|
||||
TRUSTED_CALLBACK_VARS_FIELD: callback_vars,
|
||||
"metadata": {}, # mutable-ok: Logging arg
|
||||
"model_info": {}, # mutable-ok: Logging arg
|
||||
}
|
||||
)
|
||||
return _TeamCallbackWiring(
|
||||
|
|
|
|||
|
|
@ -6,15 +6,34 @@ import json
|
|||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.models.user import LiteLLM_UserTable
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.models.user import LiteLLM_UserTable, SCIMPlaceholder
|
||||
from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
from prisma import models as prisma_models
|
||||
|
||||
_JSON_ENCODED_COLUMNS: Final = frozenset({"metadata", "model_spend", "model_max_budget"})
|
||||
|
||||
_SHADOWING_PLACEHOLDERS_SQL: Final = """
|
||||
SELECT p.user_id AS placeholder_user_id,
|
||||
array_agg(r.user_id ORDER BY r.user_id) AS resolved_user_ids,
|
||||
p.teams AS team_ids
|
||||
FROM "LiteLLM_UserTable" p
|
||||
JOIN "LiteLLM_UserTable" r
|
||||
ON r.user_id <> p.user_id
|
||||
AND (r.sso_user_id = p.user_id OR LOWER(r.user_email) = LOWER(p.user_id))
|
||||
WHERE p.sso_user_id IS NULL
|
||||
AND NOT EXISTS (SELECT 1 FROM "LiteLLM_VerificationToken" k WHERE k.user_id = p.user_id)
|
||||
GROUP BY p.user_id, p.teams
|
||||
ORDER BY p.user_id
|
||||
"""
|
||||
|
||||
_PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...])
|
||||
|
||||
|
||||
class UserRepository(BaseRepository[LiteLLM_UserTable]):
|
||||
"""Repository for user database operations."""
|
||||
|
|
@ -59,6 +78,11 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]):
|
|||
"""Find all users in a team."""
|
||||
return await self.find_many(where={"teams": {"has": team_id}})
|
||||
|
||||
async def find_shadowing_placeholders(self, tx: "Prisma") -> tuple[SCIMPlaceholder, ...]:
|
||||
"""Users with no SSO id and no virtual keys whose id is another user's SSO id or email."""
|
||||
rows: Final = await tx.query_raw(_SHADOWING_PLACEHOLDERS_SQL)
|
||||
return _PLACEHOLDER_ROWS_ADAPTER.validate_python(rows)
|
||||
|
||||
async def count_billable_users(self) -> int:
|
||||
"""Number of users that count toward the license seat limit.
|
||||
|
||||
|
|
|
|||
|
|
@ -1023,7 +1023,7 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
transformed: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> None:
|
||||
self._events: list[ResponsesAPIStreamingResponse] = _build_synthetic_response_events(
|
||||
self._events: Sequence[ResponsesAPIStreamingResponse] = build_synthetic_response_events(
|
||||
transformed=transformed,
|
||||
logging_obj=logging_obj,
|
||||
chunk_size=self.CHUNK_SIZE,
|
||||
|
|
@ -1090,7 +1090,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
transformed: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> None:
|
||||
self._events = _build_synthetic_response_events(
|
||||
self._events = build_synthetic_response_events(
|
||||
transformed=transformed,
|
||||
logging_obj=logging_obj,
|
||||
chunk_size=MockResponsesAPIStreamingIterator.CHUNK_SIZE,
|
||||
|
|
@ -1274,10 +1274,10 @@ def _add_text_like_part_events(
|
|||
)
|
||||
|
||||
|
||||
def _build_synthetic_response_events(
|
||||
def build_synthetic_response_events(
|
||||
*,
|
||||
transformed: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
logging_obj: LiteLLMLoggingObj | None,
|
||||
chunk_size: int,
|
||||
) -> list[ResponsesAPIStreamingResponse]:
|
||||
openai_types: Final = _get_openai_response_types()
|
||||
|
|
|
|||
|
|
@ -270,6 +270,10 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_deployment_rpm_limit",
|
||||
"litellm_remaining_api_key_requests_for_model",
|
||||
"litellm_remaining_api_key_tokens_for_model",
|
||||
"litellm_api_key_rate_limit_allowed_metric",
|
||||
"litellm_api_key_rate_limit_used_metric",
|
||||
"litellm_team_rate_limit_allowed_metric",
|
||||
"litellm_team_rate_limit_used_metric",
|
||||
"litellm_llm_api_failed_requests_metric",
|
||||
"litellm_callback_logging_failures_metric",
|
||||
"litellm_in_flight_requests",
|
||||
|
|
@ -775,6 +779,22 @@ class PrometheusMetricLabels:
|
|||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
]
|
||||
|
||||
litellm_api_key_rate_limit_allowed_metric: ClassVar[tuple[str, ...]] = (
|
||||
UserAPIKeyLabelNames.API_KEY_HASH.value,
|
||||
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
|
||||
UserAPIKeyLabelNames.RATE_LIMIT_TYPE.value,
|
||||
)
|
||||
|
||||
litellm_api_key_rate_limit_used_metric = litellm_api_key_rate_limit_allowed_metric
|
||||
|
||||
litellm_team_rate_limit_allowed_metric: ClassVar[tuple[str, ...]] = (
|
||||
UserAPIKeyLabelNames.TEAM.value,
|
||||
UserAPIKeyLabelNames.TEAM_ALIAS.value,
|
||||
UserAPIKeyLabelNames.RATE_LIMIT_TYPE.value,
|
||||
)
|
||||
|
||||
litellm_team_rate_limit_used_metric = litellm_team_rate_limit_allowed_metric
|
||||
|
||||
litellm_llm_api_failed_requests_metric = [
|
||||
UserAPIKeyLabelNames.END_USER.value,
|
||||
UserAPIKeyLabelNames.API_KEY_HASH.value,
|
||||
|
|
|
|||
|
|
@ -150,6 +150,12 @@ class SCIMGroup(SCIMResource):
|
|||
members: list[SCIMMember] | None = None
|
||||
|
||||
|
||||
class SCIMPlaceholderMergeResult(BaseModel):
|
||||
placeholder_user_id: str
|
||||
merged_into_user_id: str
|
||||
team_ids: tuple[str, ...]
|
||||
|
||||
|
||||
# SCIM List Response Models
|
||||
class SCIMListResponse(BaseModel):
|
||||
schemas: list[str] = ["urn:ietf:params:scim:api:messages:2.0:ListResponse"]
|
||||
|
|
|
|||
|
|
@ -3636,6 +3636,8 @@ all_litellm_params = (
|
|||
"client",
|
||||
"rpm",
|
||||
"tpm",
|
||||
"default_api_key_rpm_limit",
|
||||
"default_api_key_tpm_limit",
|
||||
"itpm",
|
||||
"otpm",
|
||||
"max_parallel_requests",
|
||||
|
|
|
|||
|
|
@ -841,7 +841,7 @@ def test_build_synthetic_response_events_covers_annotations_function_calls_and_r
|
|||
)
|
||||
|
||||
try:
|
||||
events = streaming_module._build_synthetic_response_events(
|
||||
events = streaming_module.build_synthetic_response_events(
|
||||
transformed=transformed,
|
||||
logging_obj=logging_obj,
|
||||
chunk_size=5,
|
||||
|
|
|
|||
|
|
@ -541,6 +541,74 @@ def test_llm_call_adapter_extracts_cache_tokens_from_usage_object():
|
|||
assert data.usage.cache_read_input_tokens == 3
|
||||
|
||||
|
||||
def test_llm_call_adapter_normalizes_nested_cache_tokens():
|
||||
cases: Final = (
|
||||
({"prompt_tokens_details": {"cached_tokens": 3}}, 3, None),
|
||||
({"prompt_cache_hit_tokens": 11}, 11, None),
|
||||
({"prompt_tokens_details": {"cache_write_tokens": 7}}, None, 7),
|
||||
({"prompt_tokens_details": {"cache_creation_tokens": 13}}, None, 13),
|
||||
({"prompt_tokens_details": {"cache_creation_input_tokens": 17}}, None, 17),
|
||||
)
|
||||
for usage_object, expected_read, expected_creation in cases:
|
||||
case_payload = _sample_payload(metadata={"usage_object": usage_object})
|
||||
data = LLMCallSpanData.from_standard_logging_payload(case_payload)
|
||||
assert data.usage.cache_read_input_tokens == expected_read
|
||||
assert data.usage.cache_creation_input_tokens == expected_creation
|
||||
|
||||
|
||||
def test_llm_call_adapter_prefers_nested_count_over_zero_top_level():
|
||||
payload = _sample_payload(
|
||||
metadata={
|
||||
"usage_object": {
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"prompt_tokens_details": {"cached_tokens": 5, "cache_write_tokens": 7},
|
||||
}
|
||||
}
|
||||
)
|
||||
data = LLMCallSpanData.from_standard_logging_payload(payload)
|
||||
assert data.usage.cache_read_input_tokens == 5
|
||||
assert data.usage.cache_creation_input_tokens == 7
|
||||
|
||||
|
||||
def test_llm_call_adapter_ignores_invalid_cache_values_before_valid_fallbacks():
|
||||
payload = _sample_payload(
|
||||
metadata={
|
||||
"usage_object": {
|
||||
"cache_read_input_tokens": -1,
|
||||
"cache_creation_input_tokens": "5.0",
|
||||
"prompt_tokens_details": {"cached_tokens": 5, "cache_write_tokens": 7},
|
||||
}
|
||||
}
|
||||
)
|
||||
data = LLMCallSpanData.from_standard_logging_payload(payload)
|
||||
assert data.usage.cache_read_input_tokens == 5
|
||||
assert data.usage.cache_creation_input_tokens == 7
|
||||
|
||||
|
||||
def test_llm_call_adapter_ignores_non_finite_cache_values():
|
||||
payload = _sample_payload(
|
||||
metadata={
|
||||
"usage_object": {
|
||||
"prompt_tokens_details": {"cached_tokens": float("nan")},
|
||||
}
|
||||
}
|
||||
)
|
||||
data = LLMCallSpanData.from_standard_logging_payload(payload)
|
||||
assert data.usage.cache_read_input_tokens is None
|
||||
|
||||
|
||||
def test_llm_call_adapter_preserves_explicit_zero_and_omits_missing_cache_tokens():
|
||||
for usage_object, expected_read, expected_creation in (
|
||||
({"prompt_tokens_details": {"cached_tokens": 0}}, 0, None),
|
||||
({}, None, None),
|
||||
):
|
||||
case_payload = _sample_payload(metadata={"usage_object": usage_object})
|
||||
data = LLMCallSpanData.from_standard_logging_payload(case_payload)
|
||||
assert data.usage.cache_read_input_tokens == expected_read
|
||||
assert data.usage.cache_creation_input_tokens == expected_creation
|
||||
|
||||
|
||||
def test_llm_call_adapter_cache_tokens_none_without_usage_object():
|
||||
data = LLMCallSpanData.from_standard_logging_payload(_sample_payload())
|
||||
assert data.usage.cache_creation_input_tokens is None
|
||||
|
|
|
|||
|
|
@ -93,6 +93,7 @@ async def test_async_post_call_success_hook_includes_client_ip_user_agent():
|
|||
logger._increment_token_metrics = MagicMock()
|
||||
logger._increment_remaining_budget_metrics = AsyncMock()
|
||||
logger._set_virtual_key_rate_limit_metrics = MagicMock()
|
||||
logger._set_key_and_team_rate_limit_metrics = MagicMock()
|
||||
logger._set_latency_metrics = MagicMock()
|
||||
logger.set_llm_deployment_success_metrics = MagicMock()
|
||||
logger._increment_cache_metrics = MagicMock()
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ Covers two follow-up gaps to the unified rate-limit error work:
|
|||
429s don't silently break when the new class lands.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -471,3 +472,254 @@ def test_should_ignore_non_int_v3_header_values(bad_value):
|
|||
logger.litellm_remaining_api_key_tokens_for_model.labels.return_value.set.assert_called_once_with(
|
||||
sys.maxsize
|
||||
)
|
||||
|
||||
|
||||
KEY_AND_TEAM_RATE_LIMIT_METRICS = (
|
||||
"litellm_api_key_rate_limit_allowed_metric",
|
||||
"litellm_api_key_rate_limit_used_metric",
|
||||
"litellm_team_rate_limit_allowed_metric",
|
||||
"litellm_team_rate_limit_used_metric",
|
||||
)
|
||||
|
||||
|
||||
def _clear_prometheus_registry() -> None:
|
||||
from prometheus_client import REGISTRY
|
||||
|
||||
for collector in list(REGISTRY._collector_to_names.keys()):
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _collected_samples(metric_name: str) -> dict[tuple[tuple[str, str], ...], float]:
|
||||
from prometheus_client import REGISTRY
|
||||
|
||||
return {
|
||||
tuple(sorted(sample.labels.items())): sample.value
|
||||
for metric in REGISTRY.collect()
|
||||
for sample in metric.samples
|
||||
if sample.name == metric_name
|
||||
}
|
||||
|
||||
|
||||
def _success_kwargs_with_rate_limit_headers(additional_headers: Mapping[str, object] | None) -> dict[str, object]:
|
||||
return {
|
||||
"model": "claude-haiku-4-5",
|
||||
"litellm_params": {"metadata": {}},
|
||||
"standard_logging_object": {
|
||||
"id": "t",
|
||||
"call_type": "completion",
|
||||
"response_cost": 0.001,
|
||||
"status": "success",
|
||||
"total_tokens": 20,
|
||||
"prompt_tokens": 15,
|
||||
"completion_tokens": 5,
|
||||
"startTime": 1.0,
|
||||
"endTime": 2.0,
|
||||
"completionStartTime": 1.5,
|
||||
"model": "claude-haiku-4-5",
|
||||
"model_id": "model-123",
|
||||
"model_group": "anthropic-haiku-4-5",
|
||||
"api_base": "https://api.anthropic.com",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"request_tags": [],
|
||||
"end_user": None,
|
||||
"cache_hit": False,
|
||||
"stream": False,
|
||||
"response": None,
|
||||
"model_parameters": None,
|
||||
"metadata": {
|
||||
"user_api_key_hash": "key-hash",
|
||||
"user_api_key_alias": "key-alias",
|
||||
"user_api_key_team_id": "team-id",
|
||||
"user_api_key_team_alias": "team-alias",
|
||||
"user_api_key_user_id": "u",
|
||||
"user_api_key_user_email": "e@x.com",
|
||||
"user_api_key_org_id": None,
|
||||
"user_api_key_org_alias": None,
|
||||
"requester_metadata": None,
|
||||
"user_api_key_end_user_id": None,
|
||||
"usage_object": None,
|
||||
},
|
||||
"hidden_params": {
|
||||
"litellm_overhead_time_ms": None,
|
||||
"additional_headers": additional_headers,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def _run_success_event(
|
||||
additional_headers: Mapping[str, object] | None, logger: PrometheusLogger | None = None
|
||||
) -> None:
|
||||
import datetime
|
||||
|
||||
now = datetime.datetime.now()
|
||||
await (logger or PrometheusLogger()).async_log_success_event(
|
||||
_success_kwargs_with_rate_limit_headers(additional_headers), None, now, now
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_emit_key_and_team_rate_limit_allowed_and_used_from_v3_headers():
|
||||
"""
|
||||
LIT-1672: the v3 limiter mirrors ``x-ratelimit-{api_key,team}-{limit,remaining}-*``
|
||||
into the logging payload. The gauges must expose the configured limit as-is
|
||||
and the window consumption as ``limit - remaining`` for each key / team
|
||||
dimension, split by ``rate_limit_type``.
|
||||
"""
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
await _run_success_event(
|
||||
{
|
||||
"x-ratelimit-api_key-limit-requests": 10,
|
||||
"x-ratelimit-api_key-remaining-requests": 7,
|
||||
"x-ratelimit-api_key-limit-tokens": 20000,
|
||||
"x-ratelimit-api_key-remaining-tokens": 19947,
|
||||
"x-ratelimit-team-limit-requests": 50,
|
||||
"x-ratelimit-team-remaining-requests": 47,
|
||||
"x-ratelimit-team-limit-tokens": 40000,
|
||||
"x-ratelimit-team-remaining-tokens": 39960,
|
||||
"x-ratelimit-model_per_key-limit-requests": 5,
|
||||
"x-ratelimit-model_per_key-remaining-requests": 1,
|
||||
}
|
||||
)
|
||||
|
||||
key_requests = (
|
||||
("api_key_alias", "key-alias"),
|
||||
("hashed_api_key", "key-hash"),
|
||||
("rate_limit_type", "requests"),
|
||||
)
|
||||
key_tokens = (
|
||||
("api_key_alias", "key-alias"),
|
||||
("hashed_api_key", "key-hash"),
|
||||
("rate_limit_type", "tokens"),
|
||||
)
|
||||
team_requests = (
|
||||
("rate_limit_type", "requests"),
|
||||
("team", "team-id"),
|
||||
("team_alias", "team-alias"),
|
||||
)
|
||||
team_tokens = (
|
||||
("rate_limit_type", "tokens"),
|
||||
("team", "team-id"),
|
||||
("team_alias", "team-alias"),
|
||||
)
|
||||
|
||||
assert _collected_samples("litellm_api_key_rate_limit_allowed_metric") == {
|
||||
key_requests: 10,
|
||||
key_tokens: 20000,
|
||||
}
|
||||
assert _collected_samples("litellm_api_key_rate_limit_used_metric") == {
|
||||
key_requests: 3,
|
||||
key_tokens: 53,
|
||||
}
|
||||
assert _collected_samples("litellm_team_rate_limit_allowed_metric") == {
|
||||
team_requests: 50,
|
||||
team_tokens: 40000,
|
||||
}
|
||||
assert _collected_samples("litellm_team_rate_limit_used_metric") == {
|
||||
team_requests: 3,
|
||||
team_tokens: 40,
|
||||
}
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_emit_only_the_dimensions_the_limiter_enforced():
|
||||
"""
|
||||
A key with only ``rpm_limit`` set and no team limits produces only the
|
||||
key/requests headers, so no tokens series and no team series may appear
|
||||
(a phantom 0 or sys.maxsize series would misreport an unlimited dimension).
|
||||
"""
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
await _run_success_event(
|
||||
{
|
||||
"x-ratelimit-api_key-limit-requests": 10,
|
||||
"x-ratelimit-api_key-remaining-requests": 10,
|
||||
}
|
||||
)
|
||||
|
||||
key_requests = (
|
||||
("api_key_alias", "key-alias"),
|
||||
("hashed_api_key", "key-hash"),
|
||||
("rate_limit_type", "requests"),
|
||||
)
|
||||
assert _collected_samples("litellm_api_key_rate_limit_allowed_metric") == {key_requests: 10}
|
||||
assert _collected_samples("litellm_api_key_rate_limit_used_metric") == {key_requests: 0}
|
||||
assert _collected_samples("litellm_team_rate_limit_allowed_metric") == {}
|
||||
assert _collected_samples("litellm_team_rate_limit_used_metric") == {}
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_drop_key_and_team_series_once_the_limiter_stops_reporting_a_limit():
|
||||
"""
|
||||
Removing a key's ``rpm_limit`` / ``tpm_limit`` (or a team's ``tpm_limit``)
|
||||
makes the v3 limiter stop emitting that descriptor's headers on later
|
||||
requests. The old allowed/used samples must disappear instead of keeping
|
||||
a limit that no longer exists on the scrape.
|
||||
"""
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
logger = PrometheusLogger()
|
||||
await _run_success_event(
|
||||
{
|
||||
"x-ratelimit-api_key-limit-requests": 10,
|
||||
"x-ratelimit-api_key-remaining-requests": 7,
|
||||
"x-ratelimit-api_key-limit-tokens": 20000,
|
||||
"x-ratelimit-api_key-remaining-tokens": 19947,
|
||||
"x-ratelimit-team-limit-requests": 50,
|
||||
"x-ratelimit-team-remaining-requests": 47,
|
||||
"x-ratelimit-team-limit-tokens": 40000,
|
||||
"x-ratelimit-team-remaining-tokens": 39960,
|
||||
},
|
||||
logger=logger,
|
||||
)
|
||||
await _run_success_event(
|
||||
{
|
||||
"x-ratelimit-team-limit-requests": 50,
|
||||
"x-ratelimit-team-remaining-requests": 46,
|
||||
},
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
team_requests = (
|
||||
("rate_limit_type", "requests"),
|
||||
("team", "team-id"),
|
||||
("team_alias", "team-alias"),
|
||||
)
|
||||
assert _collected_samples("litellm_api_key_rate_limit_allowed_metric") == {}
|
||||
assert _collected_samples("litellm_api_key_rate_limit_used_metric") == {}
|
||||
assert _collected_samples("litellm_team_rate_limit_allowed_metric") == {team_requests: 50}
|
||||
assert _collected_samples("litellm_team_rate_limit_used_metric") == {team_requests: 4}
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"additional_headers",
|
||||
[
|
||||
None,
|
||||
{"x-ratelimit-model_per_key-remaining-requests": 42},
|
||||
{"x-ratelimit-api_key-limit-requests": 10},
|
||||
{"x-ratelimit-api_key-limit-requests": "10", "x-ratelimit-api_key-remaining-requests": "7"},
|
||||
{"x-ratelimit-team-limit-tokens": True, "x-ratelimit-team-remaining-tokens": 5},
|
||||
],
|
||||
)
|
||||
async def test_should_emit_no_key_or_team_rate_limit_series_without_a_complete_int_pair(
|
||||
additional_headers,
|
||||
):
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
await _run_success_event(additional_headers)
|
||||
|
||||
for metric_name in KEY_AND_TEAM_RATE_LIMIT_METRICS:
|
||||
assert _collected_samples(metric_name) == {}, metric_name
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
|
|
|||
|
|
@ -3,11 +3,13 @@ import threading
|
|||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
|
|
@ -46,6 +48,43 @@ async def test_make_call_passes_logging_obj_to_client_post():
|
|||
assert call_kwargs.get("logging_obj") is logging_obj
|
||||
|
||||
|
||||
def test_anthropic_completion_does_not_send_deployment_default_limits():
|
||||
captured_requests: list[httpx.Request] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
captured_requests.append(request)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msg_default_limits",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-3-5-haiku-20241022",
|
||||
"content": [{"type": "text", "text": "Hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1},
|
||||
},
|
||||
)
|
||||
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond)))
|
||||
try:
|
||||
litellm.completion(
|
||||
model="anthropic/claude-3-5-haiku-20241022",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
api_key="test-key",
|
||||
client=client,
|
||||
default_api_key_rpm_limit=60,
|
||||
default_api_key_tpm_limit=5000000,
|
||||
)
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
request_body = json.loads(captured_requests[0].content)
|
||||
assert "default_api_key_rpm_limit" not in request_body
|
||||
assert "default_api_key_tpm_limit" not in request_body
|
||||
|
||||
|
||||
def test_redacted_thinking_content_block_delta():
|
||||
chunk = {
|
||||
"type": "content_block_start",
|
||||
|
|
|
|||
|
|
@ -1559,3 +1559,87 @@ class TestScanOnlyToolResults:
|
|||
assert data["messages"][3]["content"] == "page says [BLOCKED] here"
|
||||
assert data["messages"][3]["tool_call_id"] == "call_1"
|
||||
assert data["messages"][4]["content"] == "and then?"
|
||||
|
||||
|
||||
class TestBuildBlockSseChunks:
|
||||
"""build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE chunks"""
|
||||
|
||||
def _exc(self, original_response=None):
|
||||
from litellm.exceptions import ModifyResponseException
|
||||
|
||||
return ModifyResponseException(
|
||||
message="Blocked by policy.",
|
||||
model="gpt-5.4-mini",
|
||||
request_data={},
|
||||
guardrail_name="test",
|
||||
original_response=original_response,
|
||||
)
|
||||
|
||||
def _payloads(self, chunks):
|
||||
return [json.loads(chunk.decode().removeprefix("data: ").strip()) for chunk in chunks]
|
||||
|
||||
def test_standalone_block_uses_fresh_identity_and_zero_usage(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
first, final = self._payloads(handler.build_block_sse_chunks(self._exc(), stream_started=False))
|
||||
assert first["id"].startswith("chatcmpl-")
|
||||
assert first["model"] == "gpt-5.4-mini"
|
||||
assert first["choices"][0]["delta"] == {"role": "assistant", "content": "Blocked by policy."}
|
||||
assert first["choices"][0]["finish_reason"] is None
|
||||
assert final["choices"][0]["delta"] == {}
|
||||
assert final["choices"][0]["finish_reason"] == "content_filter"
|
||||
assert final["usage"] == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
|
||||
def test_continuation_reuses_stream_identity_and_real_usage(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
yielded = [
|
||||
{"id": "chatcmpl-live", "created": 1724900000, "model": "gpt-5.4-mini-2026-01-01"},
|
||||
]
|
||||
original = yielded + [
|
||||
{"id": "chatcmpl-live", "usage": {"prompt_tokens": 11, "completion_tokens": 5}},
|
||||
]
|
||||
first, final = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=original), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
assert (first["id"], first["created"], first["model"]) == (
|
||||
"chatcmpl-live",
|
||||
1724900000,
|
||||
"gpt-5.4-mini-2026-01-01",
|
||||
)
|
||||
assert first["choices"][0]["delta"] == {"content": "Blocked by policy."}
|
||||
assert final["id"] == "chatcmpl-live"
|
||||
assert final["usage"] == {"prompt_tokens": 11, "completion_tokens": 5, "total_tokens": 16}
|
||||
|
||||
|
||||
class TestCheckStreamingHasEnded:
|
||||
"""_check_streaming_has_ended lets end_of_stream_only withhold the finish chunk until moderation"""
|
||||
|
||||
def test_empty_and_content_only_chunks_are_not_ended(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
assert handler._check_streaming_has_ended([]) is False
|
||||
content_only = [
|
||||
{"id": "chatcmpl-live", "choices": [{"index": 0, "delta": {"content": "hi"}, "finish_reason": None}]},
|
||||
{"id": "chatcmpl-live", "choices": []},
|
||||
{"id": "chatcmpl-live", "usage": {"prompt_tokens": 1, "completion_tokens": 1}},
|
||||
]
|
||||
assert handler._check_streaming_has_ended(content_only) is False
|
||||
|
||||
def test_dict_finish_chunk_marks_stream_ended(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = [
|
||||
{"id": "chatcmpl-live", "choices": [{"index": 0, "delta": {"content": "hi"}, "finish_reason": None}]},
|
||||
{"id": "chatcmpl-live", "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]},
|
||||
]
|
||||
assert handler._check_streaming_has_ended(chunks) is True
|
||||
|
||||
def test_object_finish_chunk_marks_stream_ended(self):
|
||||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
chunks = [
|
||||
ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=None), finish_reason="stop")]
|
||||
)
|
||||
]
|
||||
assert handler._check_streaming_has_ended(chunks) is True
|
||||
|
|
|
|||
|
|
@ -1321,3 +1321,219 @@ class TestOpenAIResponsesHandlerToolInjection:
|
|||
names = [t.get("name") for t in result["tools"]]
|
||||
assert "get_weather" in names
|
||||
assert "injected_tool" in names
|
||||
|
||||
|
||||
class TestBuildBlockSseChunks:
|
||||
"""build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE events"""
|
||||
|
||||
def _exc(self, original_response=None):
|
||||
from litellm.exceptions import ModifyResponseException
|
||||
|
||||
return ModifyResponseException(
|
||||
message="Blocked by policy.",
|
||||
model="gpt-5.4-mini",
|
||||
request_data={},
|
||||
guardrail_name="test",
|
||||
original_response=original_response,
|
||||
)
|
||||
|
||||
def _payloads(self, chunks):
|
||||
import json
|
||||
|
||||
return [json.loads(chunk.decode().removeprefix("data: ").strip()) for chunk in chunks]
|
||||
|
||||
def test_standalone_block_emits_complete_synthetic_stream(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
payloads = self._payloads(handler.build_block_sse_chunks(self._exc(), stream_started=False))
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[0] == "response.created"
|
||||
assert types[-1] == "response.completed"
|
||||
completed = payloads[-1]["response"]
|
||||
assert completed["id"].startswith("resp_")
|
||||
assert completed["model"] == "gpt-5.4-mini"
|
||||
assert completed["output"][0]["content"][0]["text"] == "Blocked by policy."
|
||||
|
||||
def test_continuation_appends_item_at_next_output_index_with_real_usage(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
yielded = [
|
||||
{"type": "response.created", "response": {"id": "resp_live", "model": "gpt-5.4-mini-2026-01-01"}},
|
||||
{"type": "response.output_item.added", "output_index": 2, "item": {"id": "msg_orig"}},
|
||||
]
|
||||
original = yielded + [
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_live",
|
||||
"model": "gpt-5.4-mini-2026-01-01",
|
||||
"output": [],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28},
|
||||
},
|
||||
}
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=original), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert "response.created" not in types
|
||||
assert types[0] == "response.output_item.done"
|
||||
assert payloads[0]["output_index"] == 2
|
||||
assert payloads[0]["item"]["id"] == "msg_orig"
|
||||
assert payloads[0]["item"]["status"] == "completed"
|
||||
assert types[1] == "response.output_item.added"
|
||||
assert payloads[1]["output_index"] == 3
|
||||
completed = payloads[-1]["response"]
|
||||
assert completed["id"] == "resp_live"
|
||||
assert completed["model"] == "gpt-5.4-mini-2026-01-01"
|
||||
assert completed["output"][0]["content"][0]["text"] == "Blocked by policy."
|
||||
assert completed["usage"] == {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28}
|
||||
|
||||
def test_continuation_reads_usage_from_typed_completed_event(self):
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
original = [
|
||||
ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=ResponsesAPIResponse.model_validate(
|
||||
{
|
||||
"id": "resp_live",
|
||||
"created_at": 1,
|
||||
"model": "gpt-5.4-mini",
|
||||
"output": [],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28},
|
||||
}
|
||||
),
|
||||
)
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=original), stream_started=True, responses_so_far=[]
|
||||
)
|
||||
)
|
||||
completed = payloads[-1]["response"]
|
||||
assert completed["usage"]["input_tokens"] == 7
|
||||
assert completed["usage"]["output_tokens"] == 21
|
||||
assert completed["usage"]["total_tokens"] == 28
|
||||
|
||||
def test_continuation_closes_open_item_given_pydantic_events_with_enum_types(self):
|
||||
from litellm.types.llms.openai import (
|
||||
BaseLiteLLMOpenAIResponseObject,
|
||||
ContentPartAddedEvent,
|
||||
OutputItemAddedEvent,
|
||||
OutputTextDeltaEvent,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
open_item = GenericResponseOutputItem.model_validate(
|
||||
{"type": "message", "id": "msg_live", "status": "in_progress", "role": "assistant", "content": []}
|
||||
)
|
||||
yielded = [
|
||||
OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=0, item=open_item
|
||||
),
|
||||
ContentPartAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
|
||||
item_id="msg_live",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
part=BaseLiteLLMOpenAIResponseObject.model_validate(
|
||||
{"type": "output_text", "text": "", "annotations": []}
|
||||
),
|
||||
),
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_live",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="partial ",
|
||||
),
|
||||
OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_live",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta="text",
|
||||
),
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=yielded), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[:3] == [
|
||||
"response.output_text.done",
|
||||
"response.content_part.done",
|
||||
"response.output_item.done",
|
||||
]
|
||||
assert payloads[0]["text"] == "partial text"
|
||||
assert payloads[2]["item"]["id"] == "msg_live"
|
||||
assert payloads[2]["item"]["status"] == "completed"
|
||||
assert payloads[2]["item"]["content"][0]["text"] == "partial text"
|
||||
assert types[3] == "response.output_item.added"
|
||||
assert payloads[3]["output_index"] == 1
|
||||
|
||||
def test_continuation_closes_open_function_call_as_incomplete(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
yielded = [
|
||||
{"type": "response.created", "response": {"id": "resp_live", "model": "gpt-5.4-mini"}},
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"output_index": 0,
|
||||
"item": {
|
||||
"id": "fc_live",
|
||||
"type": "function_call",
|
||||
"status": "in_progress",
|
||||
"call_id": "call_1",
|
||||
"name": "run_payment",
|
||||
"arguments": "",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"item_id": "fc_live",
|
||||
"output_index": 0,
|
||||
"delta": '{"amount": 100}',
|
||||
},
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=yielded), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[0] == "response.output_item.done"
|
||||
closed = payloads[0]["item"]
|
||||
assert closed["id"] == "fc_live"
|
||||
assert closed["type"] == "function_call"
|
||||
assert closed["status"] == "incomplete"
|
||||
assert closed["name"] == "run_payment"
|
||||
assert "content" not in closed
|
||||
assert types[1] == "response.output_item.added"
|
||||
assert payloads[1]["output_index"] == 1
|
||||
assert types[-1] == "response.completed"
|
||||
|
||||
def test_continuation_without_open_item_emits_no_closing_events(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
yielded = [
|
||||
{"type": "response.created", "response": {"id": "resp_live", "model": "gpt-5.4-mini"}},
|
||||
{"type": "response.in_progress", "response": {"id": "resp_live"}},
|
||||
]
|
||||
payloads = self._payloads(
|
||||
handler.build_block_sse_chunks(
|
||||
self._exc(original_response=yielded), stream_started=True, responses_so_far=yielded
|
||||
)
|
||||
)
|
||||
types = [payload["type"] for payload in payloads]
|
||||
assert types[0] == "response.output_item.added"
|
||||
assert types[-1] == "response.completed"
|
||||
dones = [payload for payload in payloads if payload["type"] == "response.output_item.done"]
|
||||
assert len(dones) == 1
|
||||
assert dones[0]["item"]["content"][0]["text"] == "Blocked by policy."
|
||||
|
|
|
|||
|
|
@ -5524,7 +5524,9 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
|
|||
"""Regression for PR #38722: a topicPolicy DENY caught by the end-of-stream
|
||||
scan used to raise after SSE headers were flushed, so the client saw a
|
||||
silently truncated stream. The unified hook must emit the chat in-stream
|
||||
error frame instead."""
|
||||
error frame instead. The finish chunk is withheld while the end-of-stream
|
||||
scan runs, so on a block it is dropped rather than relayed before the
|
||||
frame."""
|
||||
from litellm.llms import load_guardrail_translation_mappings
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
|
||||
unified_guardrail as unified_module,
|
||||
|
|
@ -5582,8 +5584,9 @@ async def test_streaming_end_of_stream_block_emits_error_frame_instead_of_trunca
|
|||
finally:
|
||||
unified_module.endpoint_guardrail_translation_mappings = None
|
||||
|
||||
assert len(out) == 3
|
||||
assert len(out) == 2
|
||||
assert isinstance(out[0], ModelResponseStream)
|
||||
assert out[0].choices[0].finish_reason is None
|
||||
frame = out[-1]
|
||||
assert isinstance(frame, bytes)
|
||||
payload = json.loads(frame.decode()[len("data: ") :])
|
||||
|
|
|
|||
|
|
@ -0,0 +1,327 @@
|
|||
"""
|
||||
Regression tests for blocking an OpenAI-format streaming response from the
|
||||
unified guardrail post-call streaming iterator hook.
|
||||
|
||||
When a guardrail's ``apply_guardrail`` raises ``ModifyResponseException``
|
||||
while (or at the end of) a chat completions or Responses API stream is being
|
||||
relayed, the hook must emit a well-formed SSE termination sequence carrying
|
||||
the block message - NOT a bare ``data: {"error": ...}`` blob that surfaces as
|
||||
an HTTP 500 error frame and truncates the stream.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, AsyncGenerator, Dict, Literal, Optional, Tuple, Union
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
Delta,
|
||||
GenericGuardrailAPIInputs,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
|
||||
BLOCK_MESSAGE = "This response was replaced by policy."
|
||||
|
||||
JsonPayload = Dict[str, object]
|
||||
StreamChunk = Union[ModelResponseStream, JsonPayload, bytes]
|
||||
|
||||
|
||||
class _BlockingGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that always blocks response scans by raising ModifyResponseException."""
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
raise ModifyResponseException(
|
||||
message=BLOCK_MESSAGE,
|
||||
model="gpt-5.4-mini",
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
|
||||
|
||||
class _PassingGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that always lets response scans through unchanged."""
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return inputs
|
||||
|
||||
|
||||
def _chat_chunk(delta: Delta, finish_reason: Optional[str] = None) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-live",
|
||||
created=1724900000,
|
||||
model="gpt-5.4-mini",
|
||||
choices=[StreamingChoices(index=0, delta=delta, finish_reason=finish_reason)],
|
||||
)
|
||||
|
||||
|
||||
async def _chat_stream(end: bool) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
yield _chat_chunk(Delta(role="assistant", content="This "))
|
||||
for text in ["is ", "the ", "original ", "answer."]:
|
||||
yield _chat_chunk(Delta(content=text))
|
||||
if end:
|
||||
yield _chat_chunk(Delta(), finish_reason="stop")
|
||||
|
||||
|
||||
async def _responses_stream(end: bool) -> AsyncGenerator[JsonPayload, None]:
|
||||
original_text = "This is the original answer."
|
||||
response_envelope = {"id": "resp_live", "model": "gpt-5.4-mini", "status": "in_progress", "output": []}
|
||||
yield {"type": "response.created", "response": response_envelope}
|
||||
yield {"type": "response.in_progress", "response": response_envelope}
|
||||
yield {
|
||||
"type": "response.output_item.added",
|
||||
"output_index": 0,
|
||||
"item": {"id": "msg_orig", "type": "message", "role": "assistant", "content": []},
|
||||
}
|
||||
yield {
|
||||
"type": "response.content_part.added",
|
||||
"item_id": "msg_orig",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"part": {"type": "output_text", "text": "", "annotations": []},
|
||||
}
|
||||
for delta in ["This ", "is ", "the ", "original ", "answer."]:
|
||||
yield {
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_orig",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": delta,
|
||||
}
|
||||
yield {
|
||||
"type": "response.output_text.done",
|
||||
"item_id": "msg_orig",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"text": original_text,
|
||||
}
|
||||
if end:
|
||||
yield {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_live",
|
||||
"model": "gpt-5.4-mini",
|
||||
"status": "completed",
|
||||
"output": [
|
||||
{
|
||||
"id": "msg_orig",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": original_text, "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def _run_hook(
|
||||
route: str,
|
||||
stream: AsyncGenerator[Union[ModelResponseStream, JsonPayload], None],
|
||||
sampling_rate: int = 1,
|
||||
end_of_stream_only: bool = False,
|
||||
buffer_until_moderated: bool = False,
|
||||
blocks: bool = True,
|
||||
) -> Tuple[StreamChunk, ...]:
|
||||
guardrail = (
|
||||
_BlockingGuardrail(guardrail_name="test-blocking-guardrail", event_hook="post_call")
|
||||
if blocks
|
||||
else _PassingGuardrail(guardrail_name="test-passing-guardrail", event_hook="post_call")
|
||||
)
|
||||
guardrail.streaming_sampling_rate = sampling_rate
|
||||
guardrail.streaming_end_of_stream_only = end_of_stream_only
|
||||
guardrail.streaming_buffer_until_moderated = buffer_until_moderated
|
||||
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route=route)
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": [guardrail.guardrail_name]},
|
||||
}
|
||||
|
||||
return tuple(
|
||||
[
|
||||
chunk
|
||||
async for chunk in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=stream,
|
||||
request_data=request_data,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _sse_payloads(collected: Tuple[StreamChunk, ...]) -> Tuple[JsonPayload, ...]:
|
||||
return tuple(
|
||||
json.loads(line[len("data:") :].strip())
|
||||
for chunk in collected
|
||||
if isinstance(chunk, bytes)
|
||||
for block in chunk.decode().split("\n\n")
|
||||
for line in block.strip().split("\n")
|
||||
if line.startswith("data:")
|
||||
)
|
||||
|
||||
|
||||
def _assert_no_error_frame(collected: Tuple[StreamChunk, ...]) -> None:
|
||||
raw = "".join(chunk.decode() for chunk in collected if isinstance(chunk, bytes))
|
||||
assert '"error"' not in raw, f"unexpected error blob in stream: {raw!r}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_pre_stream_block_emits_standalone_completion():
|
||||
"""Block on the first chunk: a standalone completion opens with a role delta
|
||||
and ends with finish_reason content_filter."""
|
||||
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=False))
|
||||
_assert_no_error_frame(collected)
|
||||
payloads = _sse_payloads(collected)
|
||||
assert payloads, "no block SSE chunks were emitted"
|
||||
assert payloads[0]["choices"][0]["delta"] == {"role": "assistant", "content": BLOCK_MESSAGE}
|
||||
assert payloads[-1]["choices"][0]["finish_reason"] == "content_filter"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_mid_stream_block_continues_the_completion():
|
||||
"""Regression for the LIT-6496 500 error frame: after chunks were already
|
||||
forwarded, the block continues the same completion id and terminates with
|
||||
finish_reason content_filter instead of raising into an error blob."""
|
||||
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=False), sampling_rate=5)
|
||||
_assert_no_error_frame(collected)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, ModelResponseStream)]
|
||||
assert forwarded, "original chunks should have streamed before the block"
|
||||
payloads = _sse_payloads(collected)
|
||||
assert payloads, "no block SSE chunks were emitted"
|
||||
assert all(payload["id"] == "chatcmpl-live" for payload in payloads), (
|
||||
"block chunks must continue the in-progress completion, not start a new one"
|
||||
)
|
||||
assert payloads[0]["choices"][0]["delta"] == {"content": BLOCK_MESSAGE}
|
||||
assert payloads[-1]["choices"][0]["finish_reason"] == "content_filter"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_end_of_stream_block_terminates_cleanly():
|
||||
"""Regression for bugbot's finish-ordering finding: in end_of_stream_only
|
||||
mode the original finish chunk must be withheld until moderation decides,
|
||||
so a block's content_filter finish is the only stream terminator a client
|
||||
ever sees - never policy text trailing after finish_reason stop."""
|
||||
collected = await _run_hook("/v1/chat/completions", _chat_stream(end=True), end_of_stream_only=True)
|
||||
_assert_no_error_frame(collected)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, ModelResponseStream)]
|
||||
assert forwarded, "content chunks still stream to the client before end-of-stream moderation"
|
||||
assert all(choice.finish_reason is None for chunk in forwarded for choice in chunk.choices), (
|
||||
"the original finish chunk must be withheld until moderation decides"
|
||||
)
|
||||
payloads = _sse_payloads(collected)
|
||||
assert BLOCK_MESSAGE in json.dumps(payloads)
|
||||
assert payloads[-1]["choices"][0]["finish_reason"] == "content_filter"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_end_of_stream_pass_releases_withheld_finish_chunk():
|
||||
"""When end-of-stream moderation passes, the withheld finish chunk is
|
||||
released so a clean stream still terminates normally."""
|
||||
collected = await _run_hook(
|
||||
"/v1/chat/completions", _chat_stream(end=True), end_of_stream_only=True, blocks=False
|
||||
)
|
||||
assert not [chunk for chunk in collected if isinstance(chunk, bytes)], (
|
||||
"a clean stream must carry no synthetic block frames"
|
||||
)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, ModelResponseStream)]
|
||||
finish_reasons = [choice.finish_reason for chunk in forwarded for choice in chunk.choices]
|
||||
assert finish_reasons[-1] == "stop", "the withheld finish chunk must be released after moderation passes"
|
||||
assert all(reason is None for reason in finish_reasons[:-1])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_buffered_block_emits_full_event_sequence():
|
||||
"""Buffered moderation blocks before anything streams: a complete synthetic
|
||||
Responses stream from response.created through response.completed carrying
|
||||
the block message, with the original content never released."""
|
||||
collected = await _run_hook("/v1/responses", _responses_stream(end=True), buffer_until_moderated=True)
|
||||
_assert_no_error_frame(collected)
|
||||
assert not [chunk for chunk in collected if isinstance(chunk, dict)], (
|
||||
"buffered original chunks must never be released after a block"
|
||||
)
|
||||
payloads = _sse_payloads(collected)
|
||||
event_types = [payload["type"] for payload in payloads]
|
||||
assert event_types[0] == "response.created"
|
||||
assert "response.output_text.delta" in event_types
|
||||
assert event_types[-1] == "response.completed"
|
||||
completed = payloads[-1]["response"]
|
||||
assert completed["status"] == "completed"
|
||||
assert completed["output"][0]["content"][0]["text"] == BLOCK_MESSAGE
|
||||
assert completed["usage"] == {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28}
|
||||
assert "original answer" not in json.dumps(payloads)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_mid_stream_block_continues_the_response():
|
||||
"""Regression for the LIT-6496 500 error frame and bugbot's unclosed-item
|
||||
finding: after events were already forwarded, the block first closes the
|
||||
output item still open on the wire, then appends the replacement item under
|
||||
the same response id, and closes with response.completed - never a second
|
||||
response.created and never a completed response with an item left open."""
|
||||
collected = await _run_hook("/v1/responses", _responses_stream(end=False))
|
||||
_assert_no_error_frame(collected)
|
||||
forwarded = [chunk for chunk in collected if isinstance(chunk, dict)]
|
||||
forwarded_types = [chunk["type"] for chunk in forwarded]
|
||||
assert "response.created" in forwarded_types, "original events should have streamed before the block"
|
||||
payloads = _sse_payloads(collected)
|
||||
assert payloads, "no block SSE chunks were emitted"
|
||||
block_types = [payload["type"] for payload in payloads]
|
||||
assert "response.created" not in block_types, "a mid-stream block must not restart the response"
|
||||
assert block_types[-1] == "response.completed"
|
||||
|
||||
all_events = forwarded + list(payloads)
|
||||
opened = sorted(event["output_index"] for event in all_events if event["type"] == "response.output_item.added")
|
||||
closed = sorted(event["output_index"] for event in all_events if event["type"] == "response.output_item.done")
|
||||
assert opened == closed, "every output item opened on the stream must be closed before response.completed"
|
||||
original_done_position = block_types.index("response.output_item.done")
|
||||
block_item_position = block_types.index("response.output_item.added")
|
||||
assert original_done_position < block_item_position, (
|
||||
"the in-progress original item must be closed before the block item is appended"
|
||||
)
|
||||
assert payloads[original_done_position]["item"]["id"] == "msg_orig"
|
||||
assert payloads[block_item_position]["output_index"] == 1, (
|
||||
"the block item must continue after the original output item"
|
||||
)
|
||||
completed = payloads[-1]["response"]
|
||||
assert completed["id"] == "resp_live"
|
||||
assert completed["output"][0]["content"][0]["text"] == BLOCK_MESSAGE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_end_of_stream_block_reports_original_usage():
|
||||
collected = await _run_hook("/v1/responses", _responses_stream(end=True), end_of_stream_only=True)
|
||||
_assert_no_error_frame(collected)
|
||||
forwarded_types = [chunk["type"] for chunk in collected if isinstance(chunk, dict)]
|
||||
assert "response.completed" not in forwarded_types, (
|
||||
"the original terminal event must be withheld and replaced by the block sequence"
|
||||
)
|
||||
payloads = _sse_payloads(collected)
|
||||
completed = payloads[-1]["response"]
|
||||
assert payloads[-1]["type"] == "response.completed"
|
||||
assert completed["id"] == "resp_live"
|
||||
assert completed["output"][0]["content"][0]["text"] == BLOCK_MESSAGE
|
||||
assert completed["usage"] == {"input_tokens": 7, "output_tokens": 21, "total_tokens": 28}
|
||||
|
|
@ -1844,7 +1844,8 @@ class TestStreamingHttpErrorFrames:
|
|||
|
||||
out = await _drive_stream(UnifiedLLMGuardrails(), guardrail, chunks)
|
||||
|
||||
assert out[:2] == chunks
|
||||
assert out[0] == chunks[0]
|
||||
assert chunks[1] not in out
|
||||
frame = out[-1]
|
||||
assert isinstance(frame, bytes)
|
||||
text = frame.decode()
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, call
|
||||
|
||||
|
|
@ -38,6 +39,7 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
|
|||
get_groups,
|
||||
get_users,
|
||||
get_service_provider_config,
|
||||
merge_placeholder,
|
||||
patch_group,
|
||||
patch_team_membership,
|
||||
patch_user,
|
||||
|
|
@ -52,6 +54,7 @@ from litellm.types.proxy.management_endpoints.scim_v2 import (
|
|||
SCIMMember,
|
||||
SCIMPatchOp,
|
||||
SCIMPatchOperation,
|
||||
SCIMPlaceholderMergeResult,
|
||||
SCIMServiceProviderConfig,
|
||||
SCIMUser,
|
||||
SCIMUserEmail,
|
||||
|
|
@ -778,13 +781,17 @@ async def test_handle_existing_user_by_email_without_teams_preserves_memberships
|
|||
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
mock_team_member_add = mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the helper
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add",
|
||||
AsyncMock(),
|
||||
mock_team_member_add = (
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the helper
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add",
|
||||
AsyncMock(),
|
||||
)
|
||||
)
|
||||
mock_team_member_delete = mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the helper
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete",
|
||||
AsyncMock(),
|
||||
mock_team_member_delete = (
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the helper
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete",
|
||||
AsyncMock(),
|
||||
)
|
||||
)
|
||||
|
||||
new_user_request = NewUserRequest(
|
||||
|
|
@ -4470,9 +4477,11 @@ async def test_create_group_applies_default_team_params(
|
|||
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
|
||||
AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())),
|
||||
)
|
||||
new_team_mock = mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_group
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
|
||||
AsyncMock(return_value=mocker.MagicMock()),
|
||||
new_team_mock = (
|
||||
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_group
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
|
||||
AsyncMock(return_value=mocker.MagicMock()),
|
||||
)
|
||||
)
|
||||
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_group
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
|
||||
|
|
@ -4927,9 +4936,7 @@ async def test_process_group_patch_remove_by_the_id_the_directory_added_with(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_group_patch_remove_still_drops_a_placeholder_by_its_literal_id(
|
||||
mocker, scim_upsert_user_enabled
|
||||
):
|
||||
async def test_process_group_patch_remove_still_drops_a_placeholder_by_its_literal_id(mocker, scim_upsert_user_enabled):
|
||||
"""An earlier release put unmatched ids on the roster verbatim, so a remove has to
|
||||
keep clearing the id as written even once it also resolves."""
|
||||
patch_ops = SCIMPatchOp(
|
||||
|
|
@ -4940,7 +4947,10 @@ async def test_process_group_patch_remove_still_drops_a_placeholder_by_its_liter
|
|||
team_id="parent-group",
|
||||
team_alias="Parent Group",
|
||||
members=[],
|
||||
members_with_roles=[Member(user_id="legacy@example.com", role="user"), Member(user_id="keep-user", role="user")],
|
||||
members_with_roles=[
|
||||
Member(user_id="legacy@example.com", role="user"),
|
||||
Member(user_id="keep-user", role="user"),
|
||||
],
|
||||
)
|
||||
|
||||
_, final_members, _ = await _process_group_patch_operations(
|
||||
|
|
@ -5105,11 +5115,8 @@ async def test_process_group_patch_remove_refuses_when_two_members_share_the_id(
|
|||
assert "more than one member of this group" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_group_member_ids_exact_user_id_wins_when_it_names_nobody_else(
|
||||
mocker, scim_upsert_user_enabled
|
||||
):
|
||||
async def test_resolve_group_member_ids_exact_user_id_wins_when_it_names_nobody_else(mocker, scim_upsert_user_enabled):
|
||||
"""The canonical user id stays authoritative, including when the same account also
|
||||
holds that value as its email, which is how a SCIM-provisioned account is keyed."""
|
||||
prisma_client = _member_resolution_prisma(
|
||||
|
|
@ -5171,9 +5178,7 @@ async def test_resolve_group_member_ids_refuses_a_user_id_that_names_another_acc
|
|||
assert exc_info.value.status_code == 400
|
||||
assert "member-id" in str(exc_info.value.detail)
|
||||
create_user_mock.assert_not_called()
|
||||
assert any(
|
||||
record.levelno >= logging.WARNING and "someone-else" in record.getMessage() for record in caplog.records
|
||||
)
|
||||
assert any(record.levelno >= logging.WARNING and "someone-else" in record.getMessage() for record in caplog.records)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -5711,3 +5716,196 @@ async def test_patch_group_404s_when_team_deleted_mid_request(mocker):
|
|||
|
||||
assert exc_info.value.code == "404"
|
||||
assert f"Group not found with ID: {group_id}" in exc_info.value.message
|
||||
|
||||
|
||||
_SHADOW_MEMBER_VALUE: Final = "00u1shadow"
|
||||
_SHADOWED_ACCOUNT: Final = "real-1"
|
||||
_SHADOWED_GROUP: Final = "grp-eng"
|
||||
|
||||
|
||||
def _shadowed_tenant_rows() -> tuple[LiteLLM_UserTable, ...]:
|
||||
"""A placeholder keyed by the raw member value, and the real account that value names by SSO id."""
|
||||
return (
|
||||
LiteLLM_UserTable(user_id=_SHADOW_MEMBER_VALUE, user_email=_SHADOW_MEMBER_VALUE, teams=[_SHADOWED_GROUP]),
|
||||
LiteLLM_UserTable(user_id=_SHADOWED_ACCOUNT, user_email="alice@example.com", sso_user_id=_SHADOW_MEMBER_VALUE),
|
||||
)
|
||||
|
||||
|
||||
def _shadow_tenant_prisma(
|
||||
mocker: MockerFixture,
|
||||
*,
|
||||
rows: Sequence[LiteLLM_UserTable],
|
||||
keys_owned_by: Mapping[str, int] = MappingProxyType({}),
|
||||
) -> MagicMock:
|
||||
"""Prisma fake whose user rows are live: deleting one removes it from every later lookup."""
|
||||
users: Final[dict[str, LiteLLM_UserTable]] = {row.user_id: row for row in rows}
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id=_SHADOWED_GROUP,
|
||||
members=[_SHADOW_MEMBER_VALUE],
|
||||
members_with_roles=[Member(user_id=_SHADOW_MEMBER_VALUE, role="user")],
|
||||
metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True},
|
||||
)
|
||||
|
||||
async def find_unique(where: Mapping[str, str]) -> LiteLLM_UserTable | None:
|
||||
return users.get(where["user_id"])
|
||||
|
||||
def clause_matches(row: LiteLLM_UserTable, clause: Mapping[str, object]) -> bool:
|
||||
if "user_id" in clause:
|
||||
return row.user_id == clause["user_id"]
|
||||
if "sso_user_id" in clause:
|
||||
return row.sso_user_id == clause["sso_user_id"]
|
||||
email_filter: Final = clause["user_email"]
|
||||
assert isinstance(email_filter, dict)
|
||||
return (row.user_email or "").casefold() == str(email_filter["equals"]).casefold()
|
||||
|
||||
async def identity_rows(where: Mapping[str, object], take: int | None = None) -> tuple[LiteLLM_UserTable, ...]:
|
||||
clauses: Final = where["OR"]
|
||||
assert isinstance(clauses, list)
|
||||
matched: Final = tuple(row for row in users.values() if any(clause_matches(row, clause) for clause in clauses))
|
||||
return matched[:take] if take else matched
|
||||
|
||||
async def delete(where: Mapping[str, str]) -> LiteLLM_UserTable | None:
|
||||
return users.pop(where["user_id"], None)
|
||||
|
||||
async def keys_for(where: Mapping[str, object]) -> tuple[MagicMock, ...]:
|
||||
return tuple(mocker.MagicMock() for _ in range(keys_owned_by.get(str(where["user_id"]), 0)))
|
||||
|
||||
async def team_lookup(where: Mapping[str, str]) -> LiteLLM_TeamTable | None:
|
||||
return team if where["team_id"] == team.team_id else None
|
||||
|
||||
prisma_client = mocker.MagicMock()
|
||||
prisma_client.db = mocker.MagicMock()
|
||||
prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=find_unique)
|
||||
prisma_client.db.litellm_usertable.find_many = AsyncMock(side_effect=identity_rows)
|
||||
prisma_client.db.litellm_usertable.delete = AsyncMock(side_effect=delete)
|
||||
prisma_client.db.litellm_teamtable = mocker.MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=team_lookup)
|
||||
prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=team)
|
||||
prisma_client.db.litellm_verificationtoken = mocker.MagicMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=keys_for)
|
||||
prisma_client.db.litellm_invitationlink = mocker.MagicMock(delete_many=AsyncMock(return_value=0))
|
||||
prisma_client.db.litellm_organizationmembership = mocker.MagicMock(delete_many=AsyncMock(return_value=0))
|
||||
prisma_client.db.litellm_teammembership = mocker.MagicMock(delete_many=AsyncMock(return_value=0))
|
||||
return prisma_client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def shadowed_tenant(mocker, monkeypatch, scim_upsert_user_enabled) -> MagicMock:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
prisma_client: Final = _shadow_tenant_prisma(mocker, rows=_shadowed_tenant_rows())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
return prisma_client
|
||||
|
||||
|
||||
async def _push_shadow_member(prisma_client: MagicMock):
|
||||
return await _resolve_group_member_ids(
|
||||
members=[SCIMMember(value=_SHADOW_MEMBER_VALUE)],
|
||||
created_via="scim_group_membership",
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_merge_placeholder_hands_the_group_to_the_shadowed_account(mocker, shadowed_tenant):
|
||||
"""Every group push of the shadowing value is refused until the placeholder is folded into
|
||||
the real account; after the merge the same push resolves to that account."""
|
||||
team_member_add_mock = (
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the endpoint
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", AsyncMock()
|
||||
)
|
||||
)
|
||||
team_member_delete_mock = (
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the endpoint
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", AsyncMock()
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as before:
|
||||
await _push_shadow_member(shadowed_tenant)
|
||||
assert before.value.status_code == 400
|
||||
|
||||
result: Final = await merge_placeholder(user_id=_SHADOW_MEMBER_VALUE)
|
||||
|
||||
assert result == SCIMPlaceholderMergeResult(
|
||||
placeholder_user_id=_SHADOW_MEMBER_VALUE,
|
||||
merged_into_user_id=_SHADOWED_ACCOUNT,
|
||||
team_ids=(_SHADOWED_GROUP,),
|
||||
)
|
||||
added: Final = team_member_add_mock.call_args.kwargs["data"]
|
||||
assert (added.team_id, added.member.user_id) == (_SHADOWED_GROUP, _SHADOWED_ACCOUNT)
|
||||
dropped: Final = team_member_delete_mock.call_args.kwargs["data"]
|
||||
assert (dropped.team_id, dropped.user_id) == (_SHADOWED_GROUP, _SHADOW_MEMBER_VALUE)
|
||||
shadowed_tenant.db.litellm_teammembership.delete_many.assert_awaited_once_with(
|
||||
where={"user_id": _SHADOW_MEMBER_VALUE}
|
||||
)
|
||||
shadowed_tenant.db.litellm_usertable.delete.assert_awaited_once_with(where={"user_id": _SHADOW_MEMBER_VALUE})
|
||||
|
||||
after: Final = await _push_shadow_member(shadowed_tenant)
|
||||
assert after.all_member_ids == [_SHADOWED_ACCOUNT]
|
||||
assert after.created_users == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_merge_placeholder_keeps_the_placeholder_when_the_roster_write_fails(mocker, shadowed_tenant):
|
||||
"""If the real account cannot join the team, the placeholder stays on it, or the membership is gone
|
||||
from both accounts."""
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the endpoint
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add",
|
||||
AsyncMock(side_effect=Exception("database connection lost")),
|
||||
)
|
||||
team_member_delete_mock = (
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the endpoint
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", AsyncMock()
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException):
|
||||
await merge_placeholder(user_id=_SHADOW_MEMBER_VALUE)
|
||||
|
||||
team_member_delete_mock.assert_not_awaited()
|
||||
shadowed_tenant.db.litellm_usertable.delete.assert_not_awaited()
|
||||
assert await shadowed_tenant.db.litellm_usertable.find_unique(where={"user_id": _SHADOW_MEMBER_VALUE}) is not None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("rows", "keys_owned_by", "merged", "reason"),
|
||||
[
|
||||
pytest.param(_shadowed_tenant_rows(), {}, _SHADOWED_ACCOUNT, "SSO identity of its own", id="real-account"),
|
||||
pytest.param(
|
||||
_shadowed_tenant_rows(), {_SHADOW_MEMBER_VALUE: 2}, _SHADOW_MEMBER_VALUE, "2 virtual keys", id="owns-keys"
|
||||
),
|
||||
pytest.param(_shadowed_tenant_rows()[:1], {}, _SHADOW_MEMBER_VALUE, "shadows no account", id="names-nobody"),
|
||||
pytest.param(
|
||||
(*_shadowed_tenant_rows(), LiteLLM_UserTable(user_id="real-2", user_email=_SHADOW_MEMBER_VALUE.upper())),
|
||||
{},
|
||||
_SHADOW_MEMBER_VALUE,
|
||||
"names 2 accounts (real-1, real-2)",
|
||||
id="names-two-accounts",
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_merge_placeholder_refuses_rows_that_are_not_a_lone_placeholder(
|
||||
mocker, monkeypatch, scim_upsert_user_enabled, rows, keys_owned_by, merged, reason
|
||||
):
|
||||
"""Only a row with no SSO identity and no keys whose id names exactly one other account is folded;
|
||||
anything else could move memberships to the wrong person, so nothing is written."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
prisma_client: Final = _shadow_tenant_prisma(mocker, rows=rows, keys_owned_by=keys_owned_by)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
team_member_add_mock = (
|
||||
mocker.patch( # test-quality-ok: roster helpers are module-level, not injectable into the endpoint
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", AsyncMock()
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await merge_placeholder(user_id=merged)
|
||||
|
||||
assert int(exc_info.value.code) == 409
|
||||
assert reason in str(exc_info.value.message)
|
||||
team_member_add_mock.assert_not_awaited()
|
||||
prisma_client.db.litellm_usertable.delete.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import json
|
||||
import logging
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from io import BytesIO
|
||||
from types import SimpleNamespace
|
||||
|
|
@ -29,6 +30,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
websocket_passthrough_request,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
|
|
@ -5464,7 +5466,10 @@ def test_the_marker_check_distinguishes_the_two_route_kinds():
|
|||
assert request_dispatched_to_pass_through_endpoint(builtin) is False
|
||||
|
||||
|
||||
async def _drive_passthrough_request_and_capture_logging(user_api_key_dict: UserAPIKeyAuth) -> tuple[int, object]:
|
||||
async def _drive_passthrough_request_and_capture_logging(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
on_pre_call: Callable[[LiteLLMLoggingObj | None], None] | None = None,
|
||||
) -> tuple[int, LiteLLMLoggingObj | None]:
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
|
@ -5487,10 +5492,12 @@ async def _drive_passthrough_request_and_capture_logging(user_api_key_dict: User
|
|||
mock_request.query_params = QueryParams({})
|
||||
mock_request.body = AsyncMock(return_value=b'{"model": "gemini-2.0-flash"}')
|
||||
|
||||
captured_data: dict = {}
|
||||
captured_data: dict = {} # mutable-ok: the pre-call hook records the request data into it
|
||||
|
||||
async def capture_pre_call_hook(user_api_key_dict, data, call_type):
|
||||
captured_data.update(data)
|
||||
if on_pre_call is not None:
|
||||
on_pre_call(data.get("litellm_logging_obj"))
|
||||
return data
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
|
|
@ -5623,3 +5630,103 @@ async def test_resolve_team_callback_wiring_fails_open_on_operational_error():
|
|||
assert wiring.success_callbacks is None
|
||||
assert wiring.failure_callbacks is None
|
||||
assert wiring.logging_kwargs is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_leaves_guardrail_readable_metadata():
|
||||
"""A pre-call guardrail reads the request headers off the passthrough logging
|
||||
params without raising."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.hiddenlayer.hiddenlayer import (
|
||||
_logged_request_headers,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
team_id="test-team",
|
||||
team_metadata={
|
||||
"logging": [
|
||||
{
|
||||
"callback_name": "langfuse",
|
||||
"callback_type": "success_and_failure",
|
||||
"callback_vars": {
|
||||
"langfuse_public_key": "pk_test",
|
||||
"langfuse_secret_key": "sk_test",
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
observed: dict[str, dict[str, str] | BaseException] = {} # mutable-ok: the pre-call hook records into it
|
||||
|
||||
def read_headers_the_way_a_guardrail_does(logging_obj: LiteLLMLoggingObj | None) -> None:
|
||||
assert logging_obj is not None
|
||||
try:
|
||||
observed["headers"] = _logged_request_headers(logging_obj)
|
||||
except Exception as exc: # noqa: BLE001 - the regression is that this used to raise
|
||||
observed["headers"] = exc
|
||||
|
||||
status_code, logging_obj = await _drive_passthrough_request_and_capture_logging(
|
||||
user_api_key_dict, on_pre_call=read_headers_the_way_a_guardrail_does
|
||||
)
|
||||
|
||||
assert "headers" in observed, "the pre-call hook never ran, so nothing was observed"
|
||||
assert observed["headers"] == {}, f"guardrail header read failed: {observed['headers']!r}"
|
||||
assert status_code == 200
|
||||
assert logging_obj is not None
|
||||
assert logging_obj.dynamic_success_callbacks, "team success callbacks must stay wired"
|
||||
assert logging_obj.standard_callback_dynamic_params.get("langfuse_public_key") == "pk_test"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_leaves_cost_router_logger_working():
|
||||
"""The cost router's logger reads the deployment id off the passthrough logging
|
||||
params without raising. least_busy shares the read but swallows the exception,
|
||||
so this is the strategy where the break is observable."""
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler
|
||||
|
||||
handler = LowestCostLoggingHandler(router_cache=DualCache())
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
team_id="test-team",
|
||||
team_metadata={
|
||||
"logging": [
|
||||
{
|
||||
"callback_name": "langfuse",
|
||||
"callback_type": "success_and_failure",
|
||||
"callback_vars": {
|
||||
"langfuse_public_key": "pk_test",
|
||||
"langfuse_secret_key": "sk_test",
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
status_code, logging_obj = await _drive_passthrough_request_and_capture_logging(user_api_key_dict)
|
||||
assert status_code == 200
|
||||
assert logging_obj is not None
|
||||
|
||||
raised: list[logging.LogRecord] = [] # mutable-ok: logging.Handler records into it
|
||||
|
||||
class _RecordTracebacks(logging.Handler):
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
if record.exc_info is not None:
|
||||
raised.append(record)
|
||||
|
||||
recorder = _RecordTracebacks()
|
||||
verbose_logger.addHandler(recorder)
|
||||
try:
|
||||
await handler.async_log_success_event(
|
||||
kwargs=logging_obj.model_call_details,
|
||||
response_obj=None,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
)
|
||||
finally:
|
||||
verbose_logger.removeHandler(recorder)
|
||||
|
||||
assert not raised, f"cost router logger raised on the passthrough logging params: {raised[0].exc_info}"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22362
|
||||
"limit": 22359
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26776
|
||||
|
|
|
|||
137
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
137
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -13289,6 +13289,58 @@ export interface paths {
|
|||
patch: operations["patch_user_scim_v2_Users__user_id__patch"];
|
||||
trace?: never;
|
||||
};
|
||||
"/scim/v2/placeholders": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* List Placeholders
|
||||
* @description List user rows whose id is another account's SSO identity or email.
|
||||
*
|
||||
* An earlier release provisioned a group member it could not match as a user keyed
|
||||
* by the raw member value, and that row now shadows the account the value really
|
||||
* names, so every push of that member is refused. This lists those rows so an
|
||||
* operator can fold each one into the account it shadows with
|
||||
* ``POST /scim/v2/placeholders/{user_id}/merge``. A row that has an SSO identity of
|
||||
* its own or owns virtual keys is left out: someone uses that account.
|
||||
*/
|
||||
get: operations["list_placeholders_scim_v2_placeholders_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/scim/v2/placeholders/{user_id}/merge": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/**
|
||||
* Merge Placeholder
|
||||
* @description Fold a placeholder user into the one account its id names by SSO identity or email.
|
||||
*
|
||||
* The account is added to every team the placeholder is on, then the placeholder is
|
||||
* deleted the way ``DELETE /scim/v2/Users/{id}`` deletes a user, so the next group
|
||||
* push resolves the member value to the real account. Refused with 409 when the row
|
||||
* has an SSO identity of its own, owns virtual keys, or names no account or several.
|
||||
*/
|
||||
post: operations["merge_placeholder_scim_v2_placeholders__user_id__merge_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/search": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -34927,6 +34979,27 @@ export interface components {
|
|||
/** Value */
|
||||
value?: unknown | null;
|
||||
};
|
||||
/**
|
||||
* SCIMPlaceholder
|
||||
* @description A user row keyed by a value that names another account by SSO identity or email.
|
||||
*/
|
||||
SCIMPlaceholder: {
|
||||
/** Placeholder User Id */
|
||||
placeholder_user_id: string;
|
||||
/** Resolved User Ids */
|
||||
resolved_user_ids: string[];
|
||||
/** Team Ids */
|
||||
team_ids: string[];
|
||||
};
|
||||
/** SCIMPlaceholderMergeResult */
|
||||
SCIMPlaceholderMergeResult: {
|
||||
/** Merged Into User Id */
|
||||
merged_into_user_id: string;
|
||||
/** Placeholder User Id */
|
||||
placeholder_user_id: string;
|
||||
/** Team Ids */
|
||||
team_ids: string[];
|
||||
};
|
||||
/** SCIMServiceProviderConfig */
|
||||
SCIMServiceProviderConfig: {
|
||||
/** Authenticationschemes */
|
||||
|
|
@ -55858,6 +55931,70 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
list_placeholders_scim_v2_placeholders_get: {
|
||||
parameters: {
|
||||
query?: {
|
||||
feature?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["SCIMPlaceholder"][];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
merge_placeholder_scim_v2_placeholders__user_id__merge_post: {
|
||||
parameters: {
|
||||
query?: {
|
||||
feature?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path: {
|
||||
user_id: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["SCIMPlaceholderMergeResult"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
search_search_post: {
|
||||
parameters: {
|
||||
query?: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue