Merge pull request #41854 from BerriAI/litellm_backport_stable_batch_rc_1_102_0
Some checks failed
ai-gateway image / ai-gateway release image (push) Waiting to run
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled

fix: backport eight backport-stable fixes to rc/1.102.0 (#40596, #41046, #41086, #41171, #41178, #41283, #41495, #41689)
This commit is contained in:
Mateo Wang 2026-09-18 12:16:09 -07:00 • committed by GitHub
commit 1a993c3d28
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 1474 additions and 140 deletions

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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:

View file

@ -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.

View file

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

View file

@ -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)])

View file

@ -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:

View file

@ -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:

View file

@ -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)

View file

@ -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:

View 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.")

View file

@ -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",

View file

@ -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):

View file

@ -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

View file

@ -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(),

View file

@ -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

View file

@ -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]

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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()

View file

@ -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):

View file

@ -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"]

View file

@ -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}]

View file

@ -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.