mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
Merge pull request #41854 from BerriAI/litellm_backport_stable_batch_rc_1_102_0
fix: backport eight backport-stable fixes to rc/1.102.0 (#40596, #41046, #41086, #41171, #41178, #41283, #41495, #41689)
This commit is contained in:
commit
1a993c3d28
29 changed files with 1474 additions and 140 deletions
5
.github/workflows/test-code-quality.yml
vendored
5
.github/workflows/test-code-quality.yml
vendored
|
|
@ -178,7 +178,7 @@ jobs:
|
|||
version: "0.10.9"
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv sync --frozen --extra proxy --python 3.10
|
||||
run: uv sync --frozen --extra proxy --extra cli --python 3.10
|
||||
|
||||
- run: uv run --no-sync python --version
|
||||
|
||||
|
|
@ -187,3 +187,6 @@ jobs:
|
|||
|
||||
- name: Check litellm CLI
|
||||
run: uv run --no-sync litellm --version
|
||||
|
||||
- name: Check lite CLI
|
||||
run: uv run --no-sync lite version
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ from litellm.litellm_core_utils.logging_utils import (
|
|||
_assemble_complete_response_from_streaming_chunks,
|
||||
)
|
||||
from litellm.types.caching import CachedEmbedding
|
||||
from litellm.types.integrations.custom_logger import converted_stream_requested
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.rerank import RerankResponse
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -107,17 +108,31 @@ def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
|
|||
return "choices" in cached_result
|
||||
|
||||
|
||||
def _should_defer_streaming_cache_hit_callbacks(*, kwargs: dict[str, object]) -> bool:
|
||||
def _stream_replay_requested(kwargs: Mapping[str, object]) -> bool:
|
||||
if kwargs.get("stream", False) is True:
|
||||
return True
|
||||
return converted_stream_requested(kwargs) and not kwargs.get("_agentic_loop_depth")
|
||||
|
||||
|
||||
def _should_defer_streaming_cache_hit_callbacks(*, cached_result: object) -> bool:
|
||||
"""
|
||||
When stream=True, do not run success callbacks at cache-hit time.
|
||||
When the cache hit is replayed as a stream, do not run success callbacks at cache-hit time.
|
||||
|
||||
Cached chat/text completion replay uses CustomStreamWrapper; cached Responses
|
||||
replay uses CachedResponsesAPIStreamingIterator; cached Anthropic Messages
|
||||
replay uses CachedAnthropicMessagesStreamIterator. All invoke logging success
|
||||
handlers when the stream finishes; firing them here too would double-count
|
||||
spend and callback records.
|
||||
spend and callback records. A plain (non-stream) replay logs here, since nothing
|
||||
else will.
|
||||
"""
|
||||
return kwargs.get("stream", False) is True
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.response_cache import (
|
||||
CachedAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
return isinstance(
|
||||
cached_result, (CustomStreamWrapper, BaseResponsesAPIStreamingIterator, CachedAnthropicMessagesStreamIterator)
|
||||
)
|
||||
|
||||
|
||||
def _prompt_tokens_details_as_mapping(details: "PromptTokensDetailsWrapper") -> Mapping[str, object]:
|
||||
|
|
@ -267,7 +282,7 @@ class LLMCachingHandler:
|
|||
custom_llm_provider=kwargs.get("custom_llm_provider", None),
|
||||
args=args,
|
||||
)
|
||||
if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs):
|
||||
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):
|
||||
# LOG SUCCESS
|
||||
self._async_log_cache_hit_on_callbacks(
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -383,7 +398,7 @@ class LLMCachingHandler:
|
|||
is_async=False,
|
||||
)
|
||||
|
||||
if not _should_defer_streaming_cache_hit_callbacks(kwargs=kwargs):
|
||||
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=cached_result,
|
||||
start_time=start_time,
|
||||
|
|
@ -823,7 +838,7 @@ class LLMCachingHandler:
|
|||
if (call_type == CallTypes.acompletion.value or call_type == CallTypes.completion.value) and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
if kwargs.get("stream", False) is True:
|
||||
if _stream_replay_requested(kwargs):
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
cached_result=cached_result,
|
||||
call_type=call_type,
|
||||
|
|
@ -838,7 +853,7 @@ class LLMCachingHandler:
|
|||
if (
|
||||
call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value
|
||||
) and isinstance(cached_result, dict):
|
||||
if kwargs.get("stream", False) is True:
|
||||
if _stream_replay_requested(kwargs):
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
cached_result=cached_result,
|
||||
call_type=call_type,
|
||||
|
|
@ -893,7 +908,7 @@ class LLMCachingHandler:
|
|||
elif (call_type == "aresponses" or call_type == "responses") and isinstance(cached_result, dict):
|
||||
use_chat_completion_cache: Final = _is_chat_completion_cached_dict(cached_result)
|
||||
if use_chat_completion_cache:
|
||||
if kwargs.get("stream", False) is True:
|
||||
if _stream_replay_requested(kwargs):
|
||||
bridge_call_type: Final = (
|
||||
CallTypes.acompletion.value if call_type == "aresponses" else CallTypes.completion.value
|
||||
)
|
||||
|
|
@ -921,7 +936,7 @@ class LLMCachingHandler:
|
|||
):
|
||||
response_obj._hidden_params["cache_hit"] = True
|
||||
|
||||
if kwargs.get("stream", False) is True:
|
||||
if _stream_replay_requested(kwargs):
|
||||
cached_result = CachedResponsesAPIStreamingIterator(
|
||||
response=response_obj,
|
||||
logging_obj=logging_obj,
|
||||
|
|
|
|||
|
|
@ -633,6 +633,24 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"""Keep ``_response_ms`` / ``litellm_overhead_time_ms`` for a result that has no ``_hidden_params``."""
|
||||
self.response_timing_metrics = dict(timing_metrics) # mutable-ok: kept deep-copyable
|
||||
|
||||
def add_dynamic_callback(self, callback: CustomLogger) -> None:
|
||||
self.dynamic_input_callbacks = self._with_dynamic_callback(self.dynamic_input_callbacks, callback)
|
||||
self.dynamic_success_callbacks = self._with_dynamic_callback(self.dynamic_success_callbacks, callback)
|
||||
self.dynamic_async_success_callbacks = self._with_dynamic_callback(
|
||||
self.dynamic_async_success_callbacks, callback
|
||||
)
|
||||
self.dynamic_failure_callbacks = self._with_dynamic_callback(self.dynamic_failure_callbacks, callback)
|
||||
self.dynamic_async_failure_callbacks = self._with_dynamic_callback(
|
||||
self.dynamic_async_failure_callbacks, callback
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _with_dynamic_callback(
|
||||
callbacks: Sequence[str | Callable | CustomLogger] | None, callback: CustomLogger
|
||||
) -> list[str | Callable | CustomLogger]:
|
||||
existing: Final = tuple(callbacks or ())
|
||||
return [*existing, *(() if callback in existing else (callback,))]
|
||||
|
||||
def process_dynamic_callbacks(self):
|
||||
"""
|
||||
Initializes CustomLogger compatible callbacks in self.dynamic_* callbacks
|
||||
|
|
|
|||
|
|
@ -121,6 +121,7 @@ from litellm.repositories.table_repositories import (
|
|||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.router import Router
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
||||
|
|
@ -2368,13 +2369,6 @@ async def _backfill_null_user_email(
|
|||
return updated_row
|
||||
|
||||
|
||||
class UserNotFoundError(ValueError):
|
||||
"""The user row is provably absent, as opposed to merely unreadable, so a caller that reads a missing row as no user-level limits can key on it without also swallowing a database that would not answer."""
|
||||
|
||||
def __init__(self, user_id: str) -> None:
|
||||
super().__init__(f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call.")
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def get_user_object(
|
||||
user_id: str | None,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import os
|
|||
import tempfile
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Annotated, Final
|
||||
|
|
@ -25,7 +25,7 @@ LITELLM_PROXY_API_KEY_ENV: Final = "LITELLM_PROXY_API_KEY"
|
|||
_REJECTED_STATUSES: Final = frozenset((401, 403))
|
||||
|
||||
|
||||
class ListingFailure(StrEnum):
|
||||
class ListingFailure(str, Enum):
|
||||
"""Why a proxy could not be listed, decided once where the HTTP outcome is classified.
|
||||
|
||||
`unreachable` means no response at all; the other kinds prove the proxy answered, so callers
|
||||
|
|
|
|||
|
|
@ -32,7 +32,11 @@ from litellm.constants import (
|
|||
UNSAFE_PROXY_RESPONSE_HEADERS,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_or_create_metadata_bucket,
|
||||
independent_snapshot,
|
||||
is_expected_client_error,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
|
||||
from litellm.litellm_core_utils.get_supported_openai_params import (
|
||||
get_supported_openai_params,
|
||||
|
|
@ -2034,6 +2038,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
) -> tuple[dict, LiteLLMLoggingObj]:
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
configured_fallbacks: Final = (
|
||||
self._configured_fallbacks(llm_router=llm_router, user_api_key_dict=user_api_key_dict)
|
||||
if llm_router is not None and not self.data.get("disable_fallbacks")
|
||||
else None
|
||||
)
|
||||
pristine: Final = independent_snapshot(self.data) if configured_fallbacks else None
|
||||
|
||||
try:
|
||||
return await self.common_processing_pre_call_logic(
|
||||
request=request,
|
||||
|
|
@ -2052,14 +2063,19 @@ class ProxyBaseLLMRequestProcessing:
|
|||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyRateLimitError as original_exc:
|
||||
original_model: Final = self.data.get("model")
|
||||
if not original_model or not llm_router or self.data.get("disable_fallbacks"):
|
||||
rate_limited_data: Final = self.data
|
||||
original_model: Final = rate_limited_data.get("model")
|
||||
if (
|
||||
pristine is None
|
||||
or not configured_fallbacks
|
||||
or rate_limited_data.get("disable_fallbacks")
|
||||
or not isinstance(original_model, str)
|
||||
):
|
||||
raise
|
||||
|
||||
fallback_models: Final = self._resolve_fallback_models(
|
||||
model=original_model,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
fallbacks=configured_fallbacks,
|
||||
)
|
||||
if not fallback_models:
|
||||
raise
|
||||
|
|
@ -2074,6 +2090,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
for fallback_model in fallback_models:
|
||||
if fallback_model == original_model:
|
||||
continue
|
||||
self.data = independent_snapshot(pristine)
|
||||
self.data["model"] = fallback_model
|
||||
try:
|
||||
return await self.common_processing_pre_call_logic(
|
||||
|
|
@ -2095,39 +2112,30 @@ class ProxyBaseLLMRequestProcessing:
|
|||
except ProxyRateLimitError:
|
||||
continue
|
||||
except BaseException:
|
||||
self.data["model"] = original_model
|
||||
self.data = rate_limited_data
|
||||
raise
|
||||
|
||||
self.data["model"] = original_model
|
||||
self.data = rate_limited_data
|
||||
raise original_exc
|
||||
|
||||
def _resolve_fallback_models(
|
||||
self,
|
||||
model: str,
|
||||
llm_router: Router,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> list | None:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallbacks = None
|
||||
|
||||
@staticmethod
|
||||
def _configured_fallbacks(llm_router: Router, user_api_key_dict: UserAPIKeyAuth) -> list | None:
|
||||
key_router_settings: Final = user_api_key_dict.router_settings
|
||||
if isinstance(key_router_settings, dict) and "fallbacks" in key_router_settings:
|
||||
fallbacks = key_router_settings["fallbacks"]
|
||||
key_fallbacks: Final = key_router_settings.get("fallbacks") if isinstance(key_router_settings, dict) else None
|
||||
fallbacks: Final = key_fallbacks if key_fallbacks is not None else llm_router.fallbacks
|
||||
return fallbacks if isinstance(fallbacks, list) and fallbacks else None
|
||||
|
||||
if fallbacks is None:
|
||||
fallbacks = llm_router.fallbacks
|
||||
|
||||
if not fallbacks:
|
||||
return None
|
||||
@staticmethod
|
||||
def _resolve_fallback_models(model: str, fallbacks: list) -> list | None:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallback_model_group, generic_fallback_idx = get_fallback_model_group(
|
||||
fallbacks=fallbacks,
|
||||
model_group=model,
|
||||
)
|
||||
if fallback_model_group is None and generic_fallback_idx is not None:
|
||||
fallback_model_group = fallbacks[generic_fallback_idx]["*"]
|
||||
return fallback_model_group
|
||||
if fallback_model_group is not None:
|
||||
return fallback_model_group
|
||||
return fallbacks[generic_fallback_idx]["*"] if generic_fallback_idx is not None else None
|
||||
|
||||
@staticmethod
|
||||
def _get_model_id_from_response(hidden_params: Mapping[str, object], data: Mapping[str, object]) -> str:
|
||||
|
|
|
|||
|
|
@ -81,6 +81,11 @@ class WriterPinnedClient:
|
|||
self.db: Final = db.writer if isinstance(db, RoutingPrismaWrapper) and not db.writer_unavailable else db
|
||||
|
||||
|
||||
def writer_wrapper(db: "PrismaWrapper | RoutingPrismaWrapper") -> PrismaWrapper:
|
||||
"""Unlike `WriterPinnedClient`, ignores `writer_unavailable`: a raw SQL write has no replica fallback."""
|
||||
return db.writer if isinstance(db, RoutingPrismaWrapper) else db
|
||||
|
||||
|
||||
class RoutingPrismaWrapper:
|
||||
"""
|
||||
Routes Prisma operations between a writer and a reader Prisma client.
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ from litellm.types.llms.openai import (
|
|||
BaseLiteLLMOpenAIResponseObject,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.utils import CallTypesLiteral, LLMResponseTypes, SpecialEnums
|
||||
from litellm.types.utils import ADDRESSED_RESPONSE_ID_FIELD, CallTypesLiteral, LLMResponseTypes, SpecialEnums
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -32,7 +32,6 @@ if TYPE_CHECKING:
|
|||
_RESPONSES_API_PROVIDER_PREFIX: Final = "/openai"
|
||||
_RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"})
|
||||
|
||||
_ADDRESSED_RESPONSE_ID_KEY: Final = "_litellm_addressed_response_id"
|
||||
_UNMANAGED_RESPONSE_ID_DETAIL: Final = (
|
||||
"Forbidden. This response id was not issued by this proxy, so the proxy cannot tell who owns it. "
|
||||
"To let keys address responses this proxy did not issue, set "
|
||||
|
|
@ -132,7 +131,7 @@ class ResponsesIDSecurity(CustomLogger):
|
|||
if call_type not in responses_api_call_types:
|
||||
return None
|
||||
addressed_id_field: Final = "previous_response_id" if call_type == "aresponses" else "response_id"
|
||||
retained_id: Final = data.get(_ADDRESSED_RESPONSE_ID_KEY)
|
||||
retained_id: Final = data.get(ADDRESSED_RESPONSE_ID_FIELD)
|
||||
addressed_id: Final = (
|
||||
retained_id if isinstance(retained_id, str) and retained_id else data.get(addressed_id_field)
|
||||
)
|
||||
|
|
@ -140,7 +139,7 @@ class ResponsesIDSecurity(CustomLogger):
|
|||
return data
|
||||
authorized_id: Final = self._authorize_response_id(addressed_id, user_api_key_dict)
|
||||
data[addressed_id_field] = authorized_id
|
||||
data[_ADDRESSED_RESPONSE_ID_KEY] = addressed_id
|
||||
data[ADDRESSED_RESPONSE_ID_FIELD] = addressed_id
|
||||
return data
|
||||
|
||||
def _authorize_response_id(
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ from litellm.repositories.verification_token_repository import (
|
|||
VerificationTokenRepository,
|
||||
)
|
||||
from litellm.router import Router
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
|
|
@ -179,6 +180,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
from prisma import models as prisma_models
|
||||
from prisma import types as prisma_types
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
|
@ -4813,6 +4815,26 @@ async def _get_org_admin_org_ids(
|
|||
return org_ids if org_ids else None
|
||||
|
||||
|
||||
async def _get_user_team_ids_from_db(
|
||||
user_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[str, ...]:
|
||||
try:
|
||||
user: Final = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=True,
|
||||
)
|
||||
except UserNotFoundError:
|
||||
return ()
|
||||
return tuple(user.teams or ()) if user is not None else ()
|
||||
|
||||
|
||||
async def _build_team_list_where_conditions(
|
||||
prisma_client: PrismaClient,
|
||||
team_id: str | None,
|
||||
|
|
@ -4823,12 +4845,16 @@ async def _build_team_list_where_conditions(
|
|||
search: str | None = None,
|
||||
search_team_id_match: TeamIdSearchMatch = "exact",
|
||||
org_admin_org_ids: list[str] | None = None,
|
||||
own_team_ids: tuple[str, ...] = (),
|
||||
user_api_key_cache: UserApiKeyCache | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Build where conditions for team list query.
|
||||
|
||||
An org admin listing their own teams sees the union of the teams in the
|
||||
orgs they administer and `own_team_ids`, the teams they are a member of.
|
||||
|
||||
Returns None when the query is guaranteed to yield no results (e.g. user
|
||||
has no team memberships), allowing the caller to skip the DB round-trip.
|
||||
"""
|
||||
|
|
@ -4851,6 +4877,11 @@ async def _build_team_list_where_conditions(
|
|||
|
||||
if organization_id:
|
||||
where_conditions["organization_id"] = organization_id
|
||||
elif org_admin_org_ids is not None and own_team_ids:
|
||||
org_or_membership_scope: Final[prisma_types.LiteLLM_TeamTableWhereInput] = {
|
||||
"OR": [{"organization_id": {"in": org_admin_org_ids}}, {"team_id": {"in": list(own_team_ids)}}]
|
||||
}
|
||||
where_conditions["AND"] = [org_or_membership_scope]
|
||||
elif org_admin_org_ids is not None:
|
||||
# Org admin: always scope to their orgs, even when filtering by user_id.
|
||||
where_conditions["organization_id"] = {"in": org_admin_org_ids}
|
||||
|
|
@ -4982,66 +5013,72 @@ async def _enforce_list_team_v2_access(
|
|||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[str | None, list[str] | None]:
|
||||
) -> tuple[str | None, list[str] | None, tuple[str, ...]]:
|
||||
"""Enforce access control for list_team_v2.
|
||||
|
||||
- Proxy admins and admin viewers can query any teams.
|
||||
- Org admins can query teams within their organizations.
|
||||
- Org admins can query teams within their organizations, plus the teams
|
||||
they are a member of when listing their own teams.
|
||||
- Regular users can only query their own teams.
|
||||
|
||||
Returns the (possibly overridden) user_id and org_admin_org_ids.
|
||||
Returns the (possibly overridden) user_id, org_admin_org_ids and, for an
|
||||
org admin's own query, the caller's own team ids.
|
||||
"""
|
||||
is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
org_admin_org_ids: list[str] | None = None
|
||||
caller_user_id: Final = user_api_key_dict.user_id
|
||||
|
||||
if is_proxy_admin:
|
||||
return user_id, org_admin_org_ids
|
||||
return user_id, None, ()
|
||||
|
||||
# Always check org admin status so that even own-queries see
|
||||
# the full set of organisation teams, not just direct memberships.
|
||||
if user_api_key_dict.user_id:
|
||||
org_admin_org_ids = await _get_org_admin_org_ids(
|
||||
user_id=user_api_key_dict.user_id,
|
||||
org_admin_org_ids: Final = (
|
||||
await _get_org_admin_org_ids(
|
||||
user_id=caller_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if caller_user_id
|
||||
else None
|
||||
)
|
||||
|
||||
if org_admin_org_ids is not None:
|
||||
if caller_user_id and org_admin_org_ids is not None:
|
||||
# Org admin: validate org_id filter if provided
|
||||
if organization_id and organization_id not in org_admin_org_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "You can only view teams within your organizations."},
|
||||
)
|
||||
# When the caller is an org admin querying their own teams (or no
|
||||
# specific user), null out user_id so that
|
||||
# _build_team_list_where_conditions scopes only by organization_id
|
||||
# — org admins should see all teams in their orgs, not just teams
|
||||
# they are a direct member of. Keep user_id when the org admin
|
||||
# explicitly queries a *different* user's teams.
|
||||
if user_id is None or user_id == user_api_key_dict.user_id:
|
||||
user_id = None
|
||||
is_own_query: Final = user_id is None or user_id == caller_user_id
|
||||
own_team_ids: Final = (
|
||||
await _get_user_team_ids_from_db(
|
||||
user_id=caller_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if is_own_query
|
||||
else ()
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"list_team_v2: org admin access for user=%s, org_ids=%s, user_id_filter=%s",
|
||||
user_api_key_dict.user_id,
|
||||
_sanitize_for_log(caller_user_id),
|
||||
org_admin_org_ids,
|
||||
user_id,
|
||||
_sanitize_for_log(None if is_own_query else user_id),
|
||||
)
|
||||
else:
|
||||
# Not an org admin — fall back to standard route check
|
||||
if not allowed_route_check_inside_route(user_api_key_dict=user_api_key_dict, requested_user_id=user_id):
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": f"Only admin users can query all teams/other teams. Your user role={user_api_key_dict.user_role}"
|
||||
},
|
||||
)
|
||||
# Regular user — auto-inject caller's user_id
|
||||
if user_id is None:
|
||||
user_id = user_api_key_dict.user_id
|
||||
return None if is_own_query else user_id, org_admin_org_ids, own_team_ids
|
||||
|
||||
return user_id, org_admin_org_ids
|
||||
# Not an org admin — fall back to standard route check
|
||||
if not allowed_route_check_inside_route(user_api_key_dict=user_api_key_dict, requested_user_id=user_id):
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": f"Only admin users can query all teams/other teams. Your user role={user_api_key_dict.user_role}"
|
||||
},
|
||||
)
|
||||
# Regular user — auto-inject caller's user_id
|
||||
return user_id if user_id is not None else caller_user_id, None, ()
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -5119,7 +5156,7 @@ async def list_team_v2(
|
|||
)
|
||||
|
||||
# --- Access control ---
|
||||
user_id, org_admin_org_ids = await _enforce_list_team_v2_access(
|
||||
user_id, org_admin_org_ids, own_team_ids = await _enforce_list_team_v2_access(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
user_id=user_id,
|
||||
organization_id=organization_id,
|
||||
|
|
@ -5151,6 +5188,7 @@ async def list_team_v2(
|
|||
search=search,
|
||||
search_team_id_match=search_team_id_match,
|
||||
org_admin_org_ids=org_admin_org_ids,
|
||||
own_team_ids=own_team_ids,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -5247,17 +5285,16 @@ async def _authorize_and_filter_teams(
|
|||
|
||||
- Proxy admins: all teams (or filtered by user_id if provided).
|
||||
- Org admins: teams from their orgs (scoped to user_id if provided).
|
||||
- Own query (user_id matches caller): teams the user is a member of.
|
||||
- Own query (user_id matches caller): teams the user is a member of, across all orgs.
|
||||
- Others: 401.
|
||||
"""
|
||||
is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
is_own_query: Final = (
|
||||
user_id is not None and user_api_key_dict.user_id is not None and user_api_key_dict.user_id == user_id
|
||||
)
|
||||
allowed_org_ids: list[str] | None = None
|
||||
|
||||
if not is_proxy_admin:
|
||||
is_own_query: Final = (
|
||||
user_id is not None and user_api_key_dict.user_id is not None and user_api_key_dict.user_id == user_id
|
||||
)
|
||||
|
||||
# Check if user is an org admin (even for own queries, so they see org teams)
|
||||
if user_api_key_dict.user_id is not None:
|
||||
caller_user: Final = await get_user_object(
|
||||
|
|
@ -5284,33 +5321,30 @@ async def _authorize_and_filter_teams(
|
|||
},
|
||||
)
|
||||
|
||||
if allowed_org_ids is not None:
|
||||
# Org admin: query DB for teams in their orgs
|
||||
if allowed_org_ids is not None and not is_own_query:
|
||||
org_teams: Final = await _raw_team_db(TeamRepository(prisma_client)).find_many(
|
||||
where={"organization_id": {"in": allowed_org_ids}},
|
||||
include={"litellm_model_table": True},
|
||||
)
|
||||
if not user_id:
|
||||
return list(org_teams)
|
||||
# Filter org teams to only those where the target user is a member
|
||||
return [
|
||||
team
|
||||
for team in org_teams
|
||||
if team.members_with_roles and any(m.get("user_id") == user_id for m in team.members_with_roles)
|
||||
]
|
||||
elif user_id:
|
||||
# Regular user: fetch all and filter by membership (Prisma can't filter JSON arrays)
|
||||
response: Final = await _raw_team_db(TeamRepository(prisma_client)).find_many(
|
||||
include={"litellm_model_table": True}
|
||||
)
|
||||
return [
|
||||
team
|
||||
for team in response
|
||||
if team.members_with_roles and any(m.get("user_id") == user_id for m in team.members_with_roles)
|
||||
]
|
||||
else:
|
||||
|
||||
response: Final = await _raw_team_db(TeamRepository(prisma_client)).find_many(include={"litellm_model_table": True})
|
||||
if not user_id:
|
||||
# Proxy admin: all teams
|
||||
return list(await _raw_team_db(TeamRepository(prisma_client)).find_many(include={"litellm_model_table": True}))
|
||||
return list(response)
|
||||
|
||||
# Prisma can't filter JSON arrays, so membership is filtered in Python
|
||||
return [
|
||||
team
|
||||
for team in response
|
||||
if team.members_with_roles and any(m.get("user_id") == user_id for m in team.members_with_roles)
|
||||
]
|
||||
|
||||
|
||||
@router.get("/team/list", tags=["team management"], dependencies=[Depends(user_api_key_auth)])
|
||||
|
|
|
|||
|
|
@ -38,7 +38,7 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.auth_checks import (
|
||||
_delete_cache_access_object, # pyright: ignore[reportPrivateUsage] # the access-group endpoints reach for this same cache primitive
|
||||
)
|
||||
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
|
||||
from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper
|
||||
from litellm.repositories.table_repositories import AccessGroupRepository
|
||||
|
||||
|
||||
|
|
@ -75,7 +75,7 @@ _REPOINT_KEY_SQL: Final = (
|
|||
def _raw_executor(prisma_client: object) -> _RawExecutor:
|
||||
"""Narrow the untyped Prisma client down to the raw-query call this module makes, pinned to the writer."""
|
||||
db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
|
||||
return WriterPinnedClient(db).db # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
|
||||
return writer_wrapper(db) # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
|
||||
|
||||
|
||||
async def _invalidate_access_group_cache(access_group_id: str) -> None:
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from typing import Final, Protocol
|
|||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
|
||||
from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_caches
|
||||
from litellm.repositories.table_repositories import AccessGroupRepository
|
||||
from litellm.router import Router
|
||||
|
|
@ -56,7 +56,7 @@ _REMOVE_MODEL_NAME_SQL: Final = (
|
|||
|
||||
def _raw_executor(prisma_client: object) -> _RawExecutor:
|
||||
db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
|
||||
return WriterPinnedClient(db).db # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
|
||||
return writer_wrapper(db) # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
|
||||
|
||||
|
||||
def _config_sourced_sibling(llm_router: Router, deployment_id: str, model_id: str) -> bool:
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
|
|||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
|
||||
from litellm.types.integrations.custom_logger import converted_stream_requested
|
||||
from litellm.types.llms.openai import (
|
||||
PART_UNION_TYPES,
|
||||
ResponseAPIUsage,
|
||||
|
|
@ -611,7 +612,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
return
|
||||
|
||||
request_kwargs = getattr(caching_handler, "request_kwargs", None)
|
||||
if not _is_json_object(request_kwargs) or request_kwargs.get("stream") is not True:
|
||||
if not _is_json_object(request_kwargs):
|
||||
return
|
||||
if request_kwargs.get("stream") is not True and not converted_stream_requested(request_kwargs):
|
||||
return
|
||||
request_kwargs = request_kwargs.copy()
|
||||
preset_cache_key = getattr(caching_handler, "preset_cache_key", None)
|
||||
|
|
|
|||
|
|
@ -1625,6 +1625,24 @@ class Router:
|
|||
return
|
||||
await selector.async_pre_call_check(deployment, parent_otel_span)
|
||||
|
||||
def _bind_override_selector_to_request(
|
||||
self, strategy: str, selector: RouterStrategySelector | None, request_kwargs: Mapping[str, object] | None
|
||||
) -> None:
|
||||
if selector is None or request_kwargs is None or strategy in self._globally_registered_strategies():
|
||||
return
|
||||
logging_obj: Final = request_kwargs.get("litellm_logging_obj")
|
||||
if isinstance(logging_obj, LiteLLMLogging):
|
||||
logging_obj.add_dynamic_callback(selector)
|
||||
|
||||
def _globally_registered_strategies(self) -> frozenset[str]:
|
||||
configured: Final = (
|
||||
self.routing_strategy,
|
||||
*(group.routing_strategy for group in self._routing_groups.values()),
|
||||
)
|
||||
return frozenset(
|
||||
normalized for normalized in map(self._normalize_strategy, configured) if normalized is not None
|
||||
)
|
||||
|
||||
def _get_routing_context(
|
||||
self, model: str, request_kwargs: dict | None = None
|
||||
) -> tuple[str | None, RouterStrategySelector | None]:
|
||||
|
|
@ -1650,7 +1668,9 @@ class Router:
|
|||
override: Final = self._get_request_routing_strategy_override(request_kwargs)
|
||||
if override is not None:
|
||||
verbose_router_logger.debug("routing_group=request-override model=%s strategy=%s", model, override)
|
||||
return override, self._get_override_strategy_selector(override)
|
||||
override_selector: Final = self._get_override_strategy_selector(override)
|
||||
self._bind_override_selector_to_request(override, override_selector, request_kwargs)
|
||||
return override, override_selector
|
||||
|
||||
group_name: Final = model if self.get_routing_group(model) is not None else self._model_to_group.get(model)
|
||||
if group_name is None:
|
||||
|
|
|
|||
8
litellm/types/proxy/auth/auth_checks.py
Normal file
8
litellm/types/proxy/auth/auth_checks.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
"""Failure values raised by `litellm/proxy/auth/auth_checks.py`. Kept free of `litellm` imports so any proxy module can import them without joining the `litellm.proxy` import cycle."""
|
||||
|
||||
|
||||
class UserNotFoundError(ValueError):
|
||||
"""The user row is provably absent, as opposed to merely unreadable, so a caller that reads a missing row as no user-level limits can key on it without also swallowing a database that would not answer."""
|
||||
|
||||
def __init__(self, user_id: str) -> None:
|
||||
super().__init__(f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call.")
|
||||
|
|
@ -3706,6 +3706,8 @@ agentic_loop_internal_litellm_params: Final = [
|
|||
# the provider.
|
||||
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
|
||||
|
||||
ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id"
|
||||
|
||||
# Bedrock managed-batch deployment config, read from litellm_params by the batch and
|
||||
# files transformations. Listed for the same reason as the fields above: these sit on
|
||||
# a deployment that also serves chat, so leaking them into extra_body makes Bedrock
|
||||
|
|
@ -3720,7 +3722,7 @@ bedrock_batch_litellm_params: Final = (
|
|||
|
||||
all_litellm_params = (
|
||||
agentic_loop_internal_litellm_params
|
||||
+ [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params]
|
||||
+ [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params]
|
||||
+ [
|
||||
"metadata",
|
||||
"litellm_metadata",
|
||||
|
|
|
|||
|
|
@ -844,6 +844,39 @@ def _is_streaming_response_for_correlation(result: object) -> bool:
|
|||
return isinstance(result, CustomStreamWrapper)
|
||||
|
||||
|
||||
def _is_converted_stream_result(result: object) -> bool:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
return isinstance(result, (CustomStreamWrapper, BaseResponsesAPIStreamingIterator))
|
||||
|
||||
|
||||
async def _run_success_deployment_hook_on_converted_chat_stream(
|
||||
result: object, request_data: dict[str, object], call_type: str
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
|
||||
if not isinstance(result, CustomStreamWrapper):
|
||||
return
|
||||
completion_stream: Final = result.completion_stream
|
||||
if not isinstance(completion_stream, MockResponseIterator):
|
||||
return
|
||||
call_type_enum: Final = _CALL_TYPE_ENUM_MAP.get(call_type)
|
||||
if call_type_enum is None:
|
||||
return
|
||||
hooked: Final = await async_post_call_success_deployment_hook(
|
||||
request_data=request_data,
|
||||
response=completion_stream.model_response,
|
||||
call_type=call_type_enum,
|
||||
)
|
||||
if not isinstance(hooked, ModelResponse) or hooked is completion_stream.model_response:
|
||||
return
|
||||
result.completion_stream = MockResponseIterator( # rebind-ok: a new wrapper would drop headers and fire __del__
|
||||
model_response=hooked, json_mode=completion_stream.json_mode
|
||||
)
|
||||
|
||||
|
||||
# Runs once per call to check if the user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
|
||||
def function_setup(
|
||||
original_function: str,
|
||||
|
|
@ -1882,6 +1915,9 @@ def client(original_function):
|
|||
_caching_handler_response.cached_result is not None
|
||||
and _caching_handler_response.final_embedding_cached_response is None
|
||||
):
|
||||
if _is_converted_stream_result(_caching_handler_response.cached_result):
|
||||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
return _caching_handler_response.cached_result
|
||||
|
||||
elif _caching_handler_response.embedding_all_elements_cache_hit is True:
|
||||
|
|
@ -1939,10 +1975,14 @@ def client(original_function):
|
|||
raise
|
||||
end_time = datetime.datetime.now()
|
||||
|
||||
if _is_streaming_request(
|
||||
kwargs=kwargs,
|
||||
call_type=call_type,
|
||||
):
|
||||
streaming_requested: Final = _is_streaming_request(kwargs=kwargs, call_type=call_type)
|
||||
if streaming_requested or _is_converted_stream_result(result):
|
||||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["stream"] = True
|
||||
if not streaming_requested:
|
||||
await _run_success_deployment_hook_on_converted_chat_stream(
|
||||
result=result, request_data=kwargs, call_type=call_type
|
||||
)
|
||||
if "complete_response" in kwargs and kwargs["complete_response"] is True:
|
||||
chunks: Final = []
|
||||
for idx, chunk in enumerate(result):
|
||||
|
|
|
|||
|
|
@ -927,24 +927,14 @@ def test_sync_get_cache_defers_streaming_completion_hit_callbacks():
|
|||
|
||||
|
||||
def test_should_defer_streaming_cache_hit_callbacks_for_any_streaming_request():
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
kwargs={"stream": True},
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
kwargs={"stream": False},
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
_should_defer_streaming_cache_hit_callbacks(
|
||||
kwargs={},
|
||||
)
|
||||
is False
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
stream_replay = CustomStreamWrapper(
|
||||
completion_stream=iter(()), model="gpt-4o", logging_obj=logging_obj
|
||||
)
|
||||
assert _should_defer_streaming_cache_hit_callbacks(cached_result=stream_replay) is True
|
||||
assert _should_defer_streaming_cache_hit_callbacks(cached_result=ModelResponse()) is False
|
||||
assert _should_defer_streaming_cache_hit_callbacks(cached_result={"id": "msg_1"}) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ class TestSkipPreCallLogic:
|
|||
await processor.base_process_llm_request(
|
||||
request=MagicMock(spec=Request),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
route_type="aresponses",
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
llm_router=MagicMock(),
|
||||
|
|
|
|||
|
|
@ -693,3 +693,90 @@ async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monke
|
|||
assert handler.preset_cache_key is not None
|
||||
assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key
|
||||
assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch):
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
async def aanthropic_messages(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {
|
||||
"model": "claude-sonnet-5",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 16,
|
||||
"caching": True,
|
||||
"stream": False,
|
||||
"_websearch_interception_converted_stream": True,
|
||||
}
|
||||
cached_message = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
}
|
||||
await litellm.cache.async_add_cache(cached_message, **kwargs)
|
||||
handler = LLMCachingHandler(original_function=aanthropic_messages, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.aanthropic_messages.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
|
||||
|
||||
hit = await handler._async_get_cache(
|
||||
model="claude-sonnet-5",
|
||||
original_function=aanthropic_messages,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.aanthropic_messages.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit is not None and hit.cached_result == cached_message
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
|
||||
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agentic_loop_followup_cache_hit_with_converted_stream_marker_replays_as_plain_object(monkeypatch):
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
kwargs = {
|
||||
"model": "gpt-5.6",
|
||||
"messages": [{"role": "user", "content": "run the code"}],
|
||||
"caching": True,
|
||||
"stream": False,
|
||||
"_code_interpreter_interception_converted_stream": True,
|
||||
"_agentic_loop_depth": 1,
|
||||
}
|
||||
await litellm.cache.async_add_cache(
|
||||
litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "done"}}]), **kwargs
|
||||
)
|
||||
handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now())
|
||||
logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock()
|
||||
|
||||
hit = await handler._async_get_cache(
|
||||
model="gpt-5.6",
|
||||
original_function=acompletion,
|
||||
logging_obj=logging_obj,
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.acompletion.value,
|
||||
kwargs=kwargs,
|
||||
args=(),
|
||||
)
|
||||
|
||||
assert hit is not None and isinstance(hit.cached_result, litellm.ModelResponse)
|
||||
assert hit.cached_result.choices[0].message.content == "done"
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once()
|
||||
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True
|
||||
|
|
|
|||
|
|
@ -7004,3 +7004,21 @@ def test_get_additional_headers_survives_a_thread_growing_headers_mid_copy():
|
|||
assert copied["llm_provider-x-custom-1999"] == "1999"
|
||||
|
||||
_run_while_a_thread_grows(headers, read, reads=300)
|
||||
|
||||
|
||||
def test_add_dynamic_callback_registers_once_per_list_without_touching_the_callers_list(logging_obj: LitellmLogging):
|
||||
callback: Final = CustomLogger()
|
||||
caller_owned: Final = ["langfuse"]
|
||||
logging_obj.dynamic_success_callbacks = caller_owned
|
||||
|
||||
logging_obj.add_dynamic_callback(callback)
|
||||
logging_obj.add_dynamic_callback(callback)
|
||||
|
||||
assert caller_owned == ["langfuse"]
|
||||
assert logging_obj.dynamic_success_callbacks == ["langfuse", callback]
|
||||
assert logging_obj.dynamic_input_callbacks == [callback]
|
||||
assert logging_obj.dynamic_async_success_callbacks == [callback]
|
||||
assert logging_obj.dynamic_failure_callbacks == [callback]
|
||||
assert logging_obj.dynamic_async_failure_callbacks == [callback]
|
||||
assert LitellmLogging._with_dynamic_callback(None, callback) == [callback]
|
||||
assert LitellmLogging._with_dynamic_callback((callback,), callback) == [callback]
|
||||
|
|
|
|||
|
|
@ -20,6 +20,13 @@ from litellm.proxy.client.cli.commands.pi import (
|
|||
)
|
||||
|
||||
|
||||
def test_listing_failure_is_str_enum():
|
||||
assert issubclass(ListingFailure, str)
|
||||
assert ListingFailure.REJECTED.value == "rejected"
|
||||
assert ListingFailure("rejected") is ListingFailure.REJECTED
|
||||
assert str(ListingFailure.REJECTED.value) == "rejected"
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, status_code, payload=None):
|
||||
self.status_code = status_code
|
||||
|
|
|
|||
|
|
@ -145,6 +145,18 @@ def test_writer_pinned_client_yields_to_routed_reads_when_writer_down():
|
|||
assert pinned.db.litellm_proxymodeltable.find_many is reader_inner.litellm_proxymodeltable.find_many
|
||||
|
||||
|
||||
def test_writer_wrapper_keeps_raw_sql_on_the_writer_while_writer_flagged_down():
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper, writer_wrapper
|
||||
|
||||
writer, writer_inner, reader, reader_inner = _make_wrappers()
|
||||
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
routing._writer_unavailable = True
|
||||
|
||||
assert writer_wrapper(routing).query_raw is writer_inner.query_raw
|
||||
assert writer_wrapper(routing).query_raw is not reader_inner.query_raw
|
||||
assert writer_wrapper(writer) is writer
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_invokes_both_clients():
|
||||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||||
|
|
|
|||
|
|
@ -3937,6 +3937,7 @@ async def test_list_team_v2_org_admin_sees_org_teams():
|
|||
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team])
|
||||
mock_db.litellm_teamtable.count = AsyncMock(return_value=1)
|
||||
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
|
||||
mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user)
|
||||
|
||||
result = await list_team_v2(
|
||||
http_request=mock_request,
|
||||
|
|
@ -4036,6 +4037,7 @@ async def test_list_team_v2_org_admin_own_user_id_sees_all_org_teams():
|
|||
)
|
||||
mock_db.litellm_teamtable.count = AsyncMock(return_value=2)
|
||||
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
|
||||
mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user)
|
||||
|
||||
# UI sends the caller's own user_id for non-Admin roles
|
||||
result = await list_team_v2(
|
||||
|
|
@ -4055,10 +4057,217 @@ async def test_list_team_v2_org_admin_own_user_id_sees_all_org_teams():
|
|||
assert result["total"] == 2
|
||||
assert len(result["teams"]) == 2
|
||||
|
||||
# Verify the where clause scopes by org only — no team_id filter
|
||||
# Verify the where clause scopes by org OR own membership — no
|
||||
# top-level team_id filter that would hide org teams they aren't in
|
||||
where = mock_db.litellm_teamtable.find_many.call_args.kwargs["where"]
|
||||
assert where["organization_id"] == {"in": ["org_A"]}
|
||||
assert where["AND"] == [
|
||||
{"OR": [{"organization_id": {"in": ["org_A"]}}, {"team_id": {"in": ["team_1"]}}]}
|
||||
]
|
||||
assert "team_id" not in where
|
||||
assert "organization_id" not in where
|
||||
|
||||
|
||||
def _team_where_matches(team, where) -> bool:
|
||||
for key, cond in where.items():
|
||||
if key == "AND":
|
||||
if not all(_team_where_matches(team, c) for c in cond):
|
||||
return False
|
||||
elif key == "OR":
|
||||
if not any(_team_where_matches(team, c) for c in cond):
|
||||
return False
|
||||
else:
|
||||
value = getattr(team, key)
|
||||
if not isinstance(cond, dict):
|
||||
if value != cond:
|
||||
return False
|
||||
elif "in" in cond and value not in cond["in"]:
|
||||
return False
|
||||
elif "contains" in cond and cond["contains"].lower() not in (value or "").lower():
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _org_membership(user_id: str, organization_id: str, user_role: str) -> LiteLLM_OrganizationMembershipTable:
|
||||
return LiteLLM_OrganizationMembershipTable(
|
||||
user_id=user_id,
|
||||
organization_id=organization_id,
|
||||
user_role=user_role,
|
||||
spend=0.0,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs(monkeypatch):
|
||||
"""
|
||||
/v2/team/list: an org admin of org_A who is a member of a team in org_B
|
||||
gets that team back on a self query (with and without user_id, with and
|
||||
without search), alongside every org_A team. The membership half of the
|
||||
union comes from the DB, so a stale cached user object cannot hide it.
|
||||
A query for another user stays scoped to org_A.
|
||||
|
||||
Regression test for LIT-3723.
|
||||
"""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.team_endpoints import list_team_v2
|
||||
|
||||
org_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="org_admin_user")
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="org_admin_user",
|
||||
value=LiteLLM_UserTable(
|
||||
user_id="org_admin_user",
|
||||
teams=["team_in_org_A"],
|
||||
organization_memberships=[
|
||||
_org_membership("org_admin_user", "org_A", "org_admin"),
|
||||
_org_membership("org_admin_user", "org_B", "internal_user"),
|
||||
],
|
||||
),
|
||||
model_type=LiteLLM_UserTable,
|
||||
)
|
||||
await cache.async_set_cache(
|
||||
key="other_user",
|
||||
value=LiteLLM_UserTable(
|
||||
user_id="other_user",
|
||||
teams=["other_team_in_org_A", "team_in_org_B", "unrelated_team_in_org_B"],
|
||||
organization_memberships=[_org_membership("other_user", "org_B", "internal_user")],
|
||||
),
|
||||
model_type=LiteLLM_UserTable,
|
||||
)
|
||||
|
||||
def team(team_id, organization_id, *member_ids):
|
||||
return LiteLLM_TeamTable(
|
||||
team_id=team_id,
|
||||
team_alias=team_id,
|
||||
organization_id=organization_id,
|
||||
members_with_roles=[Member(user_id=m, role="user") for m in member_ids],
|
||||
)
|
||||
|
||||
all_teams = [
|
||||
team("team_in_org_A", "org_A", "org_admin_user"),
|
||||
team("other_team_in_org_A", "org_A", "other_user"),
|
||||
team("team_in_org_B", "org_B", "org_admin_user", "other_user"),
|
||||
team("unrelated_team_in_org_B", "org_B", "other_user"),
|
||||
]
|
||||
|
||||
async def find_many(where=None, **kwargs):
|
||||
return [t for t in all_teams if where is None or _team_where_matches(t, where)]
|
||||
|
||||
async def count(where=None, **kwargs):
|
||||
return len(await find_many(where))
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
|
||||
prisma_client.db.litellm_teamtable.count = AsyncMock(side_effect=count)
|
||||
prisma_client.db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(
|
||||
user_id="org_admin_user",
|
||||
teams=["team_in_org_A", "team_in_org_B"],
|
||||
organization_memberships=[
|
||||
_org_membership("org_admin_user", "org_A", "org_admin"),
|
||||
_org_membership("org_admin_user", "org_B", "internal_user"),
|
||||
],
|
||||
)
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj)
|
||||
|
||||
async def list_teams(user_id, search=None):
|
||||
result = await list_team_v2(
|
||||
http_request=MagicMock(),
|
||||
user_id=user_id,
|
||||
organization_id=None,
|
||||
team_id=None,
|
||||
team_alias=None,
|
||||
search=search,
|
||||
user_api_key_dict=org_admin,
|
||||
page=1,
|
||||
page_size=10,
|
||||
sort_by=None,
|
||||
sort_order="asc",
|
||||
status=None,
|
||||
)
|
||||
assert result["total"] == len(result["teams"])
|
||||
return [t.team_id for t in result["teams"]]
|
||||
|
||||
own_view = ["team_in_org_A", "other_team_in_org_A", "team_in_org_B"]
|
||||
assert await list_teams("org_admin_user") == own_view
|
||||
assert await list_teams(None) == own_view
|
||||
assert await list_teams("org_admin_user", search="team_in_org_B") == ["team_in_org_B"]
|
||||
assert await list_teams("other_user") == ["other_team_in_org_A"]
|
||||
prisma_client.db.litellm_usertable.find_unique.assert_awaited_with(
|
||||
where={"user_id": "org_admin_user"}, include={"organization_memberships": True}
|
||||
)
|
||||
|
||||
prisma_client.db.litellm_usertable.find_unique.side_effect = RuntimeError("db down")
|
||||
with pytest.raises(ValueError, match="db down"):
|
||||
await list_teams("org_admin_user")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_team_v1_org_admin_own_query_keeps_memberships_in_other_orgs():
|
||||
"""
|
||||
/team/list: an org admin of org_A listing their own teams sees every team
|
||||
they belong to, including the org_B one. The bare admin listing stays the
|
||||
org_A view and a query for another user stays scoped to org_A.
|
||||
|
||||
Regression test for LIT-3723.
|
||||
"""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.team_endpoints import _authorize_and_filter_teams
|
||||
|
||||
org_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="org_admin_user")
|
||||
cache = UserApiKeyCache()
|
||||
await cache.async_set_cache(
|
||||
key="org_admin_user",
|
||||
value=LiteLLM_UserTable(
|
||||
user_id="org_admin_user",
|
||||
teams=["team_in_org_A", "team_in_org_B"],
|
||||
organization_memberships=[_org_membership("org_admin_user", "org_A", "org_admin")],
|
||||
),
|
||||
model_type=LiteLLM_UserTable,
|
||||
)
|
||||
|
||||
def team(team_id, organization_id, *member_ids):
|
||||
return SimpleNamespace(
|
||||
team_id=team_id,
|
||||
organization_id=organization_id,
|
||||
members_with_roles=[{"user_id": m, "role": "user"} for m in member_ids],
|
||||
)
|
||||
|
||||
all_teams = [
|
||||
team("team_in_org_A", "org_A", "org_admin_user"),
|
||||
team("other_team_in_org_A", "org_A", "other_user"),
|
||||
team("team_in_org_B", "org_B", "org_admin_user", "other_user"),
|
||||
team("unrelated_team_in_org_B", "org_B", "other_user"),
|
||||
]
|
||||
|
||||
async def find_many(where=None, **kwargs):
|
||||
if where is None:
|
||||
return all_teams
|
||||
return [t for t in all_teams if t.organization_id in where["organization_id"]["in"]]
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
|
||||
|
||||
async def list_teams(user_id):
|
||||
teams = await _authorize_and_filter_teams(
|
||||
user_api_key_dict=org_admin,
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=cache,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
return [t.team_id for t in teams]
|
||||
|
||||
assert await list_teams("org_admin_user") == ["team_in_org_A", "team_in_org_B"]
|
||||
assert await list_teams(None) == ["team_in_org_A", "other_team_in_org_A"]
|
||||
assert await list_teams("other_user") == ["other_team_in_org_A"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -11,14 +11,15 @@ from litellm.proxy.management_helpers.access_group_key_sync import (
|
|||
)
|
||||
|
||||
|
||||
def _routed_prisma_client():
|
||||
def _routed_prisma_client(writer_unavailable: bool = False):
|
||||
writer_inner = MagicMock(name="writer_prisma")
|
||||
reader_inner = MagicMock(name="reader_prisma")
|
||||
writer_inner.query_raw = AsyncMock(return_value=[])
|
||||
reader_inner.query_raw = AsyncMock(return_value=[])
|
||||
reader_inner.query_raw = AsyncMock(side_effect=RuntimeError("cannot execute UPDATE in a read-only transaction"))
|
||||
writer = PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False)
|
||||
reader = PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False)
|
||||
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
routing._writer_unavailable = writer_unavailable
|
||||
return SimpleNamespace(db=routing), writer_inner, reader_inner
|
||||
|
||||
|
||||
|
|
@ -39,6 +40,39 @@ async def test_regeneration_repoint_update_runs_on_the_writer():
|
|||
reader_inner.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_regeneration_repoint_update_stays_on_the_writer_while_writer_flagged_unavailable():
|
||||
prisma_client, writer_inner, reader_inner = _routed_prisma_client(writer_unavailable=True)
|
||||
|
||||
await sync_key_regeneration_access_group_membership(
|
||||
prisma_client=prisma_client,
|
||||
previous_key_token="old-token",
|
||||
new_key_token="new-token",
|
||||
data=None,
|
||||
existing_key_row=MagicMock(),
|
||||
)
|
||||
|
||||
writer_inner.query_raw.assert_awaited_once()
|
||||
assert writer_inner.query_raw.await_args.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"')
|
||||
assert writer_inner.query_raw.await_args.args[1:] == ("old-token", "new-token")
|
||||
reader_inner.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_membership_attach_and_detach_updates_stay_on_the_writer_while_writer_flagged_unavailable():
|
||||
prisma_client, writer_inner, reader_inner = _routed_prisma_client(writer_unavailable=True)
|
||||
|
||||
await sync_key_access_group_membership(
|
||||
prisma_client=prisma_client,
|
||||
key_token="token",
|
||||
previous_access_group_ids=["ag-old"],
|
||||
updated_access_group_ids=["ag-new"],
|
||||
)
|
||||
|
||||
assert writer_inner.query_raw.await_count == 2
|
||||
reader_inner.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_membership_attach_and_detach_updates_run_on_the_writer():
|
||||
prisma_client, writer_inner, reader_inner = _routed_prisma_client()
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from litellm.proxy.management_helpers.access_group_model_sync import (
|
|||
_INVALIDATE = "litellm.proxy.management_helpers.access_group_model_sync.invalidate_access_group_caches"
|
||||
|
||||
|
||||
def _routed_prisma_client(deployment_count: int):
|
||||
def _routed_prisma_client(deployment_count: int, writer_unavailable: bool = False):
|
||||
async def query_raw(sql, *params):
|
||||
if sql.startswith("SELECT COUNT(*)"):
|
||||
return [{"deployment_count": deployment_count}]
|
||||
|
|
@ -26,6 +26,7 @@ def _routed_prisma_client(deployment_count: int):
|
|||
writer = PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False)
|
||||
reader = PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False)
|
||||
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
routing._writer_unavailable = writer_unavailable
|
||||
return SimpleNamespace(db=routing), writer_inner, reader_inner
|
||||
|
||||
|
||||
|
|
@ -53,6 +54,20 @@ async def test_rename_replaces_the_old_name_when_no_other_deployment_carries_it(
|
|||
reader_inner.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rename_update_stays_on_the_writer_while_writer_flagged_unavailable():
|
||||
prisma_client, writer_inner, reader_inner = _routed_prisma_client(deployment_count=0, writer_unavailable=True)
|
||||
|
||||
with patch(_INVALIDATE, new=AsyncMock()):
|
||||
await sync_access_groups_for_renamed_model(
|
||||
prisma_client, model_id="m-1", old_name="gpt-5.6", new_name="gpt-5.6-eu", llm_router=None
|
||||
)
|
||||
|
||||
(update_call,) = _access_group_updates(writer_inner)
|
||||
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
|
||||
reader_inner.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rename_appends_the_new_name_when_a_sibling_row_keeps_the_old_one():
|
||||
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=1)
|
||||
|
|
@ -168,3 +183,17 @@ async def test_delete_keeps_the_name_while_a_sibling_row_still_backs_it():
|
|||
|
||||
assert _access_group_updates(writer_inner) == []
|
||||
invalidate.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_counts_backing_rows_on_the_writer_not_a_lagging_replica_while_writer_flagged_unavailable():
|
||||
prisma_client, writer_inner, reader_inner = _routed_prisma_client(deployment_count=0, writer_unavailable=True)
|
||||
reader_inner.query_raw = AsyncMock(return_value=[{"deployment_count": 1}])
|
||||
|
||||
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
|
||||
await sync_access_groups_for_deleted_model(prisma_client, model_id="m-1", model_name="gpt-5.6", llm_router=None)
|
||||
|
||||
(update_call,) = _access_group_updates(writer_inner)
|
||||
assert "array_remove" in update_call.args[0]
|
||||
invalidate.assert_awaited_once_with(("ag-1", "ag-2"))
|
||||
reader_inner.query_raw.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.proxy.common_request_processing import (
|
|||
create_response,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -6235,6 +6236,206 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
|
|||
call_type="acompletion",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _v3_limiter_rig(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth,
|
||||
fallbacks: list[dict[str, list[str]]],
|
||||
) -> tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]]:
|
||||
"""Real v3 limiter (the default ``parallel_request_limiter``) wired in through the
|
||||
``proxy_logging_obj`` seam, so ``common_processing_pre_call_logic`` runs for real:
|
||||
``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
limiter_models: list[str] = []
|
||||
|
||||
async def run_limiter(
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth, data: dict[str, object], call_type: str
|
||||
) -> dict[str, object]:
|
||||
limiter_models.append(str(data["model"]))
|
||||
await limiter.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
)
|
||||
return data
|
||||
|
||||
proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=run_limiter)
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": group, "litellm_params": {"model": "openai/gpt-4.1-nano", "api_key": "fake"}}
|
||||
for chain in fallbacks
|
||||
for group in (*chain.keys(), *(m for models in chain.values() for m in models))
|
||||
],
|
||||
fallbacks=fallbacks,
|
||||
)
|
||||
return proxy_logging_obj, router, proxy_server.ProxyConfig(), limiter_models
|
||||
|
||||
@staticmethod
|
||||
def _otel_key(
|
||||
rpm_limit: int | None = None,
|
||||
model_rpm_limit: dict[str, int] | None = None,
|
||||
disable_fallbacks: bool = False,
|
||||
) -> ProxyUserAPIKeyAuth:
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
span = TracerProvider().get_tracer("test").start_span("proxy-request")
|
||||
return ProxyUserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
parent_otel_span=span,
|
||||
rpm_limit=rpm_limit,
|
||||
metadata={
|
||||
**({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}),
|
||||
**({"disable_fallbacks": True} if disable_fallbacks else {}),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _chat_request() -> Request:
|
||||
return Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []})
|
||||
|
||||
async def _pre_call(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth,
|
||||
rig: tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]],
|
||||
) -> tuple[ProxyBaseLLMRequestProcessing, tuple[dict[str, object], LiteLLMLoggingObj]]:
|
||||
proxy_logging_obj, router, proxy_config, _ = rig
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
result = await processor._pre_call_with_fallbacks(
|
||||
request=self._chat_request(),
|
||||
general_settings={},
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
version=None,
|
||||
proxy_config=proxy_config,
|
||||
user_model=None,
|
||||
user_temperature=None,
|
||||
user_request_timeout=None,
|
||||
user_max_tokens=None,
|
||||
user_api_base=None,
|
||||
model=None,
|
||||
route_type="acompletion",
|
||||
llm_router=router,
|
||||
)
|
||||
return processor, result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v3_limiter_with_otel_span_falls_back_from_client_request(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Customer path: OTel on, per-key model RPM cap on the primary, a router fallback configured.
|
||||
The first pass enriches ``data["metadata"]`` with the live span, then the limiter raises. The
|
||||
fallback pass must start from the client's request again, so ``add_litellm_data_to_request``
|
||||
never deep-copies the span (the ``cannot pickle '_thread.RLock'`` 500)."""
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
key = self._otel_key(model_rpm_limit={primary_model: 1})
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
|
||||
def client_request() -> dict[str, object]:
|
||||
return {
|
||||
"model": primary_model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {"tags": ["client-tag"]},
|
||||
}
|
||||
|
||||
_, (first_data, _) = await self._pre_call(client_request(), key, rig)
|
||||
processor, (data, logging_obj) = await self._pre_call(client_request(), key, rig)
|
||||
|
||||
assert first_data["model"] == primary_model
|
||||
assert data["model"] == fallback_model
|
||||
assert processor.data is data
|
||||
assert data["litellm_logging_obj"] is logging_obj
|
||||
assert logging_obj.model == fallback_model
|
||||
requester_metadata = data["metadata"]["requester_metadata"]
|
||||
assert requester_metadata["tags"] == ["client-tag"]
|
||||
assert "litellm_parent_otel_span" not in requester_metadata
|
||||
assert "user_api_key_auth" not in requester_metadata
|
||||
assert data["metadata"]["litellm_parent_otel_span"] is key.parent_otel_span
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v3_limiter_with_otel_span_returns_429_when_fallbacks_exhausted(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
key = self._otel_key(rpm_limit=1)
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
request = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
processor = ProxyBaseLLMRequestProcessing(data=dict(request))
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
await processor._pre_call_with_fallbacks(
|
||||
request=self._chat_request(),
|
||||
general_settings={},
|
||||
proxy_logging_obj=rig[0],
|
||||
user_api_key_dict=key,
|
||||
version=None,
|
||||
proxy_config=rig[2],
|
||||
user_model=None,
|
||||
user_temperature=None,
|
||||
user_request_timeout=None,
|
||||
user_max_tokens=None,
|
||||
user_api_base=None,
|
||||
model=None,
|
||||
route_type="acompletion",
|
||||
llm_router=rig[1],
|
||||
)
|
||||
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
assert exc_info.value.status_code == 429
|
||||
assert "Rate limit exceeded" in str(exc_info.value.detail)
|
||||
assert exc_info.value.headers["retry-after"]
|
||||
assert processor.data["model"] == primary_model
|
||||
assert processor.data["litellm_logging_obj"].model == primary_model
|
||||
assert processor.data["litellm_call_id"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_lookup_uses_alias_resolved_model_group(self, monkeypatch: pytest.MonkeyPatch):
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
monkeypatch.setattr(litellm, "model_alias_map", {"my-alias": primary_model})
|
||||
key = self._otel_key(model_rpm_limit={primary_model: 1})
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
request = {"model": "my-alias", "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
_, (data, _) = await self._pre_call(dict(request), key, rig)
|
||||
|
||||
assert data["model"] == fallback_model
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_metadata_disable_fallbacks_returns_429_instead_of_retrying(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""``disable_fallbacks`` set in key metadata only lands on ``data`` during the first
|
||||
pre-call pass (``add_key_level_controls``), so it must be honored after that pass."""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
primary_model = "gpt-4.1"
|
||||
fallback_model = "gpt-4.1-mini"
|
||||
key = self._otel_key(model_rpm_limit={primary_model: 1}, disable_fallbacks=True)
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
request = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert rig[3] == [primary_model, primary_model]
|
||||
|
||||
|
||||
class _RecordingSuccessLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
|
|
|
|||
|
|
@ -12,14 +12,21 @@ capture the forwarded kwargs; if the flag-setting line is removed the captured
|
|||
kwargs lack the flag and these tests fail.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.responses.litellm_completion_transformation.handler import (
|
||||
LiteLLMCompletionTransformationHandler,
|
||||
)
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.utils import ADDRESSED_RESPONSE_ID_FIELD
|
||||
|
||||
|
||||
class _StopForwarding(Exception):
|
||||
|
|
@ -68,3 +75,49 @@ async def test_async_fallback_tags_skip_responses_api_bridge():
|
|||
await coro
|
||||
|
||||
assert captured.get("_skip_responses_api_bridge") is True
|
||||
|
||||
|
||||
class _RecordingAnthropicHandler:
|
||||
def __init__(self, reply: Mapping[str, object]) -> None:
|
||||
self.reply: Final = reply
|
||||
self.request_body: Mapping[str, object] | None = None
|
||||
|
||||
def __call__(self, request: httpx.Request) -> httpx.Response:
|
||||
self.request_body = json.loads(request.content)
|
||||
return httpx.Response(200, json=dict(self.reply), request=request)
|
||||
|
||||
|
||||
_ANTHROPIC_MESSAGE_PAYLOAD: Final = {
|
||||
"id": "msg_turn_two",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-6",
|
||||
"content": [{"type": "text", "text": "14"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 12, "output_tokens": 1},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridged_follow_up_turn_keeps_the_addressed_response_id_off_the_provider_body():
|
||||
provider: Final = _RecordingAnthropicHandler(_ANTHROPIC_MESSAGE_PAYLOAD)
|
||||
client: Final = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(provider))
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model="azure_ai/claude-sonnet-4-6",
|
||||
api_base="https://fake-foundry-resource.services.ai.azure.com",
|
||||
api_key="fake-api-key",
|
||||
input="Double it",
|
||||
previous_response_id="resp_turn_one",
|
||||
client=client,
|
||||
**{ADDRESSED_RESPONSE_ID_FIELD: "resp_turn_one"},
|
||||
)
|
||||
|
||||
assert provider.request_body is not None, "the bridged turn never reached the provider"
|
||||
assert ADDRESSED_RESPONSE_ID_FIELD not in provider.request_body, (
|
||||
f"the addressed response id reached the provider body: {sorted(provider.request_body)}"
|
||||
)
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert [item.type for item in response.output] == ["message"]
|
||||
|
|
|
|||
|
|
@ -5,15 +5,20 @@ the implicit `"default"` group driven by the router's top-level
|
|||
`routing_strategy` / `routing_strategy_args`.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.router import RoutingGroup, RoutingStrategy
|
||||
from litellm.utils import Rules, function_setup
|
||||
|
||||
|
||||
def _model_list():
|
||||
|
|
@ -954,6 +959,223 @@ def test_sync_pass_through_specific_deployment_runs_the_override_pre_call_check(
|
|||
assert plain["model_info"]["id"] == "deploy-3"
|
||||
|
||||
|
||||
def _two_deployment_model_list(**d1_params: object) -> list[dict[str, object]]:
|
||||
return [
|
||||
{
|
||||
"model_name": "grp",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test-1", "mock_response": "ok", **d1_params},
|
||||
"model_info": {"id": "d1"},
|
||||
},
|
||||
{
|
||||
"model_name": "grp",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test-2", "mock_response": "ok"},
|
||||
"model_info": {"id": "d2"},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _proxy_shaped_request(**data: object) -> dict[str, object]:
|
||||
"""The proxy builds the request's `Logging` object before it hands the call to the router."""
|
||||
logging_obj, kwargs = function_setup(
|
||||
"acompletion",
|
||||
Rules(),
|
||||
datetime.datetime.now(),
|
||||
litellm_call_id=str(uuid.uuid4()),
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
**data,
|
||||
)
|
||||
return {**kwargs, "litellm_logging_obj": logging_obj}
|
||||
|
||||
|
||||
async def _async_override_pick(router: Router, strategy: str) -> str:
|
||||
deployment = await router.async_get_available_deployment(
|
||||
"grp", request_kwargs=_proxy_shaped_request(model="grp", routing_strategy=strategy)
|
||||
)
|
||||
return deployment["model_info"]["id"]
|
||||
|
||||
|
||||
def _sync_override_pick(router: Router, strategy: str) -> str:
|
||||
deployment = router.get_available_deployment(
|
||||
"grp", request_kwargs=_proxy_shaped_request(model="grp", routing_strategy=strategy)
|
||||
)
|
||||
return deployment["model_info"]["id"]
|
||||
|
||||
|
||||
def _in_flight(router: Router, deployment_id: str) -> int | None:
|
||||
return router.cache.get_cache(f"grp_request_count:{deployment_id}")
|
||||
|
||||
|
||||
async def _async_wait_until(predicate: Callable[[], bool]) -> None:
|
||||
for _ in range(100):
|
||||
if predicate():
|
||||
return
|
||||
await asyncio.sleep(0.02)
|
||||
raise AssertionError("lifecycle callback never reached the override selector")
|
||||
|
||||
|
||||
def _sync_wait_until(predicate: Callable[[], bool]) -> None:
|
||||
for _ in range(100):
|
||||
if predicate():
|
||||
return
|
||||
time.sleep(0.02)
|
||||
raise AssertionError("lifecycle callback never reached the override selector")
|
||||
|
||||
|
||||
def _selector_is_not_global(selector: CustomLogger) -> bool:
|
||||
global_lists = (
|
||||
litellm.callbacks,
|
||||
litellm.input_callback,
|
||||
litellm.success_callback,
|
||||
litellm.failure_callback,
|
||||
litellm._async_success_callback,
|
||||
litellm._async_failure_callback,
|
||||
)
|
||||
return not any(cb is selector for cbs in global_lists for cb in cbs)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_least_busy_override_sees_the_overriding_request_in_flight():
|
||||
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle", num_retries=0)
|
||||
|
||||
stream = await router.acompletion(**_proxy_shaped_request(model="grp", routing_strategy="least-busy", stream=True))
|
||||
busy = stream._hidden_params["model_id"]
|
||||
idle = "d2" if busy == "d1" else "d1"
|
||||
assert [await _async_override_pick(router, "least-busy") for _ in range(3)] == [idle, idle, idle]
|
||||
|
||||
async for _ in stream:
|
||||
pass
|
||||
await _async_wait_until(lambda: _in_flight(router, busy) == 0)
|
||||
assert await _async_override_pick(router, "least-busy") == "d1"
|
||||
assert _selector_is_not_global(router._override_selectors["least-busy"])
|
||||
|
||||
|
||||
def test_sync_least_busy_override_sees_the_overriding_request_in_flight():
|
||||
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle", num_retries=0)
|
||||
|
||||
stream = router.completion(**_proxy_shaped_request(model="grp", routing_strategy="least-busy", stream=True))
|
||||
busy = stream._hidden_params["model_id"]
|
||||
idle = "d2" if busy == "d1" else "d1"
|
||||
assert [_sync_override_pick(router, "least-busy") for _ in range(3)] == [idle, idle, idle]
|
||||
|
||||
for _ in stream:
|
||||
pass
|
||||
_sync_wait_until(lambda: _in_flight(router, busy) == 0)
|
||||
assert _sync_override_pick(router, "least-busy") == "d1"
|
||||
assert _selector_is_not_global(router._override_selectors["least-busy"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_least_busy_override_releases_the_slot_when_the_overriding_request_fails():
|
||||
router = Router(
|
||||
model_list=_two_deployment_model_list(mock_response="litellm.InternalServerError"),
|
||||
routing_strategy="simple-shuffle",
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.InternalServerError):
|
||||
await router.acompletion(**_proxy_shaped_request(model="grp", routing_strategy="least-busy"))
|
||||
|
||||
await _async_wait_until(lambda: _in_flight(router, "d1") == 0)
|
||||
assert await _async_override_pick(router, "least-busy") == "d1"
|
||||
assert _selector_is_not_global(router._override_selectors["least-busy"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_latency_based_override_learns_from_the_overriding_requests():
|
||||
router = Router(
|
||||
model_list=_two_deployment_model_list(mock_delay=0.05), routing_strategy="simple-shuffle", num_retries=0
|
||||
)
|
||||
|
||||
def samples(deployment_id: str) -> list[float]:
|
||||
recorded = (router.cache.get_cache("grp_map") or {}).get(deployment_id, {}).get("latency", [])
|
||||
return [latency for latency in recorded if latency > 0]
|
||||
|
||||
async def overriding_call() -> str:
|
||||
sampled_before = {"d1": len(samples("d1")), "d2": len(samples("d2"))}
|
||||
response = await router.acompletion(
|
||||
**_proxy_shaped_request(model="grp", routing_strategy="latency-based-routing")
|
||||
)
|
||||
deployment_id = response._hidden_params["model_id"]
|
||||
await _async_wait_until(lambda: len(samples(deployment_id)) > sampled_before[deployment_id])
|
||||
return deployment_id
|
||||
|
||||
served = [await overriding_call() for _ in range(6)]
|
||||
|
||||
assert "d1" in served
|
||||
assert served[2:] == ["d2"] * 4
|
||||
assert _selector_is_not_global(router._override_selectors["latency-based-routing"])
|
||||
|
||||
|
||||
def test_override_selector_is_bound_only_to_the_request_that_asked_for_it():
|
||||
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle")
|
||||
overriding = _proxy_shaped_request(model="grp", routing_strategy="least-busy")
|
||||
plain = _proxy_shaped_request(model="grp")
|
||||
|
||||
router.get_available_deployment("grp", request_kwargs=overriding)
|
||||
router.get_available_deployment("grp", request_kwargs=overriding)
|
||||
router.get_available_deployment("grp", request_kwargs=plain)
|
||||
|
||||
selector = router._override_selectors["least-busy"]
|
||||
bound = overriding["litellm_logging_obj"]
|
||||
for callbacks in (
|
||||
bound.dynamic_input_callbacks,
|
||||
bound.dynamic_success_callbacks,
|
||||
bound.dynamic_async_success_callbacks,
|
||||
bound.dynamic_failure_callbacks,
|
||||
bound.dynamic_async_failure_callbacks,
|
||||
):
|
||||
assert callbacks == [selector]
|
||||
unbound = plain["litellm_logging_obj"]
|
||||
assert unbound.dynamic_input_callbacks is None and unbound.dynamic_success_callbacks is None
|
||||
assert unbound.dynamic_failure_callbacks is None and unbound.dynamic_async_failure_callbacks is None
|
||||
|
||||
|
||||
def test_override_matching_the_router_strategy_is_not_bound_twice():
|
||||
router = Router(model_list=_two_deployment_model_list(), routing_strategy="least-busy")
|
||||
request = _proxy_shaped_request(model="grp", routing_strategy="least-busy")
|
||||
|
||||
router.get_available_deployment("grp", request_kwargs=request)
|
||||
|
||||
assert request["litellm_logging_obj"].dynamic_input_callbacks is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_override_matching_a_routing_group_strategy_records_each_request_once():
|
||||
router = Router(
|
||||
model_list=_two_deployment_model_list(),
|
||||
routing_strategy="simple-shuffle",
|
||||
routing_groups=[RoutingGroup(group_name="lat", models=["grp"], routing_strategy="latency-based-routing")],
|
||||
num_retries=0,
|
||||
)
|
||||
request = _proxy_shaped_request(model="grp", routing_strategy="latency-based-routing")
|
||||
assert router._globally_registered_strategies() == {"simple-shuffle", "latency-based-routing"}
|
||||
|
||||
response = await router.acompletion(**request)
|
||||
deployment_id = response._hidden_params["model_id"]
|
||||
await _async_wait_until(lambda: (router.cache.get_cache("grp_map") or {}).get(deployment_id) is not None)
|
||||
|
||||
assert len(router.cache.get_cache("grp_map")[deployment_id]["latency"]) == 1
|
||||
assert request["litellm_logging_obj"].dynamic_success_callbacks is None
|
||||
|
||||
|
||||
def test_bind_override_selector_to_request_binds_once_and_ignores_requests_without_logging():
|
||||
router = Router(model_list=_two_deployment_model_list(), routing_strategy="simple-shuffle")
|
||||
selector = router._get_override_strategy_selector("least-busy")
|
||||
request = _proxy_shaped_request(model="grp", routing_strategy="least-busy")
|
||||
request["litellm_logging_obj"].dynamic_success_callbacks = ["langfuse"]
|
||||
|
||||
router._bind_override_selector_to_request("least-busy", selector, request)
|
||||
router._bind_override_selector_to_request("least-busy", selector, request)
|
||||
router._bind_override_selector_to_request("least-busy", selector, None)
|
||||
router._bind_override_selector_to_request("least-busy", selector, {"model": "grp"})
|
||||
|
||||
logging_obj = request["litellm_logging_obj"]
|
||||
assert logging_obj.dynamic_success_callbacks == ["langfuse", selector]
|
||||
assert logging_obj.dynamic_input_callbacks == [selector]
|
||||
assert logging_obj.dynamic_async_failure_callbacks == [selector]
|
||||
assert _selector_is_not_global(selector)
|
||||
|
||||
|
||||
def _quality_group(strategy="latency-based-routing"):
|
||||
return [{"group_name": "quality", "models": ["filtered-model", "other-model"], "routing_strategy": strategy}]
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from jsonschema import validate
|
|||
|
||||
import litellm
|
||||
from litellm._internal_context import is_internal_call
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
|
||||
from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
|
||||
from litellm._logging import (
|
||||
CorrelationContextFilter,
|
||||
|
|
@ -30,24 +32,31 @@ from litellm._logging import (
|
|||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.proxy.utils import is_valid_api_key
|
||||
from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
Choices,
|
||||
Delta,
|
||||
LlmProviders,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
PromptTokensDetailsWrapper,
|
||||
StreamingChoices,
|
||||
Usage,
|
||||
ADDRESSED_RESPONSE_ID_FIELD,
|
||||
)
|
||||
from litellm.types.utils import all_litellm_params, bedrock_batch_litellm_params
|
||||
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams
|
||||
from litellm.utils import (
|
||||
CustomStreamWrapper,
|
||||
ProviderConfigManager,
|
||||
TextCompletionStreamWrapper,
|
||||
_check_provider_match,
|
||||
_get_potential_model_names,
|
||||
_is_streaming_request,
|
||||
_run_success_deployment_hook_on_converted_chat_stream,
|
||||
_snapshot_exception_for_hook,
|
||||
async_post_call_failure_deployment_hook,
|
||||
async_post_call_success_deployment_hook,
|
||||
|
|
@ -5170,6 +5179,306 @@ async def test_wrapper_async_restores_originating_task_context_after_success(mon
|
|||
session_id_var.set("")
|
||||
|
||||
|
||||
class _ConvertStreamDeploymentHook(CustomLogger):
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: dict[str, object], call_type: CallTypes | None
|
||||
) -> dict[str, object] | None:
|
||||
if not kwargs.get("stream"):
|
||||
return None
|
||||
return {**kwargs, "stream": False, HEADROOM_CONVERTED_STREAM_KEY: True}
|
||||
|
||||
|
||||
class _SuccessKwargsCapture(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.success_kwargs: list[dict[str, object]] = []
|
||||
self.stream_event_responses: list[object] = []
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.success_kwargs.append(kwargs)
|
||||
|
||||
async def async_log_stream_event(
|
||||
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.stream_event_responses.append(response_obj)
|
||||
|
||||
|
||||
def _install_converted_stream_callbacks(monkeypatch: pytest.MonkeyPatch) -> _SuccessKwargsCapture:
|
||||
capture: Final = _SuccessKwargsCapture()
|
||||
monkeypatch.setattr(litellm, "callbacks", [_ConvertStreamDeploymentHook(), capture])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
return capture
|
||||
|
||||
|
||||
async def _wait_for_success_kwargs(capture: _SuccessKwargsCapture, count: int = 1) -> dict[str, object]:
|
||||
for _ in range(50):
|
||||
if len(capture.success_kwargs) >= count and not _PENDING_CACHE_WRITES:
|
||||
break
|
||||
await asyncio.sleep(0.05)
|
||||
await asyncio.sleep(0.2)
|
||||
assert len(capture.success_kwargs) == count
|
||||
return capture.success_kwargs[-1]
|
||||
|
||||
|
||||
def _assert_cache_hit_logged_as_stream(capture: _SuccessKwargsCapture, success_kwargs: dict[str, object]) -> None:
|
||||
standard_logging_object: Final = success_kwargs["standard_logging_object"]
|
||||
assert isinstance(standard_logging_object, dict)
|
||||
assert standard_logging_object["cache_hit"] is True
|
||||
assert standard_logging_object["stream"] is True
|
||||
assert success_kwargs["stream"] is True
|
||||
assert capture.stream_event_responses == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_logs_converted_chat_stream_with_standard_logging_object(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
capture: Final = _install_converted_stream_callbacks(monkeypatch)
|
||||
|
||||
response: Final = await litellm.acompletion(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
mock_response="converted stream body",
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
chunks: Final = [chunk async for chunk in response]
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "converted stream body"
|
||||
|
||||
success_kwargs: Final = await _wait_for_success_kwargs(capture)
|
||||
standard_logging_object: Final = success_kwargs["standard_logging_object"]
|
||||
assert isinstance(standard_logging_object, dict)
|
||||
assert standard_logging_object["response_cost"] > 0
|
||||
assert standard_logging_object["stream"] is True
|
||||
assert success_kwargs["stream"] is True
|
||||
|
||||
|
||||
class _RewritingSuccessDeploymentHook(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.seen_responses: tuple[object, ...] = ()
|
||||
|
||||
async def async_post_call_success_deployment_hook(
|
||||
self, request_data: dict[str, object], response: object, call_type: CallTypes | None
|
||||
) -> ModelResponse | None:
|
||||
self.seen_responses = (*self.seen_responses, response)
|
||||
if not isinstance(response, ModelResponse):
|
||||
return None
|
||||
choice: Final = response.choices[0]
|
||||
if not isinstance(choice, Choices):
|
||||
return None
|
||||
rewritten_message: Final = choice.message.model_copy(update={"content": "rewritten by deployment hook"})
|
||||
return response.model_copy(update={"choices": [choice.model_copy(update={"message": rewritten_message})]})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_runs_success_deployment_hook_on_converted_chat_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_install_converted_stream_callbacks(monkeypatch)
|
||||
hook: Final = _RewritingSuccessDeploymentHook()
|
||||
monkeypatch.setattr(litellm, "callbacks", [_ConvertStreamDeploymentHook(), hook])
|
||||
|
||||
response: Final = await litellm.acompletion(
|
||||
model="gpt-5.6",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
mock_response="converted stream body",
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
chunks: Final = [chunk async for chunk in response]
|
||||
|
||||
assert len(hook.seen_responses) == 1
|
||||
seen: Final = hook.seen_responses[0]
|
||||
assert isinstance(seen, ModelResponse)
|
||||
assert seen.choices[0].message.content == "converted stream body"
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "rewritten by deployment hook"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("completion_stream", "call_type"),
|
||||
[
|
||||
(iter([ModelResponse(model="gpt-5.6")]), "acompletion"),
|
||||
(MockResponseIterator(model_response=ModelResponse(model="gpt-5.6")), "not_a_call_type"),
|
||||
],
|
||||
ids=["real_provider_stream", "unmapped_call_type"],
|
||||
)
|
||||
async def test_converted_chat_stream_hook_skips_unhandled_wrappers(
|
||||
monkeypatch: pytest.MonkeyPatch, completion_stream: object, call_type: str
|
||||
) -> None:
|
||||
hook: Final = _RewritingSuccessDeploymentHook()
|
||||
monkeypatch.setattr(litellm, "callbacks", [hook])
|
||||
wrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream, model="gpt-5.6", logging_obj=MagicMock(), custom_llm_provider="openai"
|
||||
)
|
||||
|
||||
await _run_success_deployment_hook_on_converted_chat_stream(
|
||||
result=wrapper, request_data={"model": "gpt-5.6"}, call_type=call_type
|
||||
)
|
||||
|
||||
assert hook.seen_responses == ()
|
||||
assert wrapper.completion_stream is completion_stream
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_wrapper_async_leaves_success_deployment_hook_off_requested_fake_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
hook: Final = _RewritingSuccessDeploymentHook()
|
||||
monkeypatch.setattr(litellm, "callbacks", [hook])
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
respx.post("http://fake-stream.invalid/api/v1/run/flow-1").respond(
|
||||
json={"outputs": [{"outputs": [{"results": {"message": {"text": "plain stream body"}}}]}]}
|
||||
)
|
||||
|
||||
response: Final = await litellm.acompletion(
|
||||
model="langflow/flow-1",
|
||||
api_base="http://fake-stream.invalid",
|
||||
api_key="fake-key",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
num_retries=0,
|
||||
)
|
||||
assert isinstance(response, CustomStreamWrapper)
|
||||
assert isinstance(response.completion_stream, MockResponseIterator)
|
||||
chunks: Final = [chunk async for chunk in response]
|
||||
|
||||
assert hook.seen_responses == ()
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "plain stream body"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_wrapper_async_logs_converted_responses_stream_with_standard_logging_object(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
capture: Final = _install_converted_stream_callbacks(monkeypatch)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
respx.post("https://api.openai.com/v1/responses").respond(
|
||||
json={
|
||||
"id": "resp_converted",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.6",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_converted",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "converted stream body", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7},
|
||||
}
|
||||
)
|
||||
|
||||
response: Final = await litellm.aresponses(
|
||||
model="openai/gpt-5.6", input="hi", stream=True, api_key="sk-test", num_retries=0
|
||||
)
|
||||
assert isinstance(response, BaseResponsesAPIStreamingIterator)
|
||||
events: Final = [event async for event in response]
|
||||
assert events[-1].type == "response.completed"
|
||||
|
||||
success_kwargs: Final = await _wait_for_success_kwargs(capture)
|
||||
standard_logging_object: Final = success_kwargs["standard_logging_object"]
|
||||
assert isinstance(standard_logging_object, dict)
|
||||
assert standard_logging_object["response_cost"] > 0
|
||||
assert standard_logging_object["stream"] is True
|
||||
assert success_kwargs["stream"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_async_replays_cached_converted_chat_stream_as_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
capture: Final = _install_converted_stream_callbacks(monkeypatch)
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
request: Final = {
|
||||
"model": "gpt-5.6",
|
||||
"messages": [{"role": "user", "content": "replay me from cache"}],
|
||||
"stream": True,
|
||||
"mock_response": "converted stream body",
|
||||
"num_retries": 0,
|
||||
}
|
||||
|
||||
first: Final = await litellm.acompletion(**request)
|
||||
first_chunks: Final = [chunk async for chunk in first]
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in first_chunks) == "converted stream body"
|
||||
await _wait_for_success_kwargs(capture)
|
||||
|
||||
replay: Final = await litellm.acompletion(**request)
|
||||
assert isinstance(replay, CustomStreamWrapper)
|
||||
replay_chunks: Final = [chunk async for chunk in replay]
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in replay_chunks) == "converted stream body"
|
||||
|
||||
_assert_cache_hit_logged_as_stream(capture, await _wait_for_success_kwargs(capture, count=2))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@respx.mock
|
||||
async def test_wrapper_async_replays_cached_converted_responses_stream_as_stream(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
|
||||
capture: Final = _install_converted_stream_callbacks(monkeypatch)
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
route: Final = respx.post("https://api.openai.com/v1/responses").respond(
|
||||
json={
|
||||
"id": "resp_cached_converted",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.6",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_cached_converted",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "converted stream body", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 3, "output_tokens": 4, "total_tokens": 7},
|
||||
}
|
||||
)
|
||||
request: Final = {
|
||||
"model": "openai/gpt-5.6",
|
||||
"input": "replay me from cache",
|
||||
"stream": True,
|
||||
"api_key": "sk-test",
|
||||
"num_retries": 0,
|
||||
}
|
||||
|
||||
first: Final = await litellm.aresponses(**request)
|
||||
assert [event async for event in first][-1].type == "response.completed"
|
||||
await _wait_for_success_kwargs(capture)
|
||||
|
||||
replay: Final = await litellm.aresponses(**request)
|
||||
assert isinstance(replay, BaseResponsesAPIStreamingIterator)
|
||||
assert [event async for event in replay][-1].type == "response.completed"
|
||||
assert route.call_count == 1
|
||||
|
||||
_assert_cache_hit_logged_as_stream(capture, await _wait_for_success_kwargs(capture, count=2))
|
||||
|
||||
|
||||
def test_function_setup_failure_after_logging_construction_restores_context(monkeypatch):
|
||||
"""If function_setup() constructs Logging() (which already mutated
|
||||
trace_id_var/session_id_var in __init__) but then raises before returning,
|
||||
|
|
@ -5295,6 +5604,20 @@ def test_get_litellm_params_keys_never_reach_the_provider():
|
|||
)
|
||||
|
||||
|
||||
def test_addressed_response_id_never_reaches_the_provider():
|
||||
kwargs = {
|
||||
"a_real_provider_specific_param": 1,
|
||||
ADDRESSED_RESPONSE_ID_FIELD: "resp_addressed-by-the-client",
|
||||
}
|
||||
|
||||
non_default = get_non_default_completion_params(kwargs)
|
||||
|
||||
assert non_default == {"a_real_provider_specific_param": 1}, (
|
||||
"the addressed response id leaked into the provider params: "
|
||||
f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}"
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_batch_params_never_reach_the_provider():
|
||||
"""A Bedrock managed-batch deployment carries aws_batch_role_arn / s3_* /
|
||||
bedrock_tags in its litellm_params, and the same deployment also serves chat.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue