mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
Merge pull request #29460 from BerriAI/litellm_cherrypick_1_85_x
Cherry-pick #29311, #29343, #29358, #27913, and #29447 onto stable/1.85.x
This commit is contained in:
commit
6a026ee617
21 changed files with 1387 additions and 365 deletions
|
|
@ -1603,6 +1603,90 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
) -> Optional[float]:
|
||||
return self._response_cost_calculator(result=result, cache_hit=cache_hit)
|
||||
|
||||
@staticmethod
|
||||
def _is_sync_litellm_request(litellm_params: dict) -> bool:
|
||||
"""True for sync SDK entrypoints (``completion``), false for async (``acompletion``, etc.)."""
|
||||
return (
|
||||
litellm_params.get(CallTypes.acompletion.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aresponses.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aembedding.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
|
||||
and litellm_params.get(CallTypes.atranscription.value, False) is not True
|
||||
)
|
||||
|
||||
def _is_assembled_stream_success(self, result=None) -> bool:
|
||||
"""Final assembled stream export (not a per-chunk success call).
|
||||
|
||||
Per-chunk callers pass a ``ModelResponseStream`` (or ``None``); the
|
||||
final assembled response is any other non-``None`` value (typically a
|
||||
``ModelResponse``). Treating a chunk as the assembled response would
|
||||
prematurely set the ``has_dispatched_final_stream_success`` dedup
|
||||
guard and silently suppress the real final stream log.
|
||||
"""
|
||||
if self.stream is not True:
|
||||
return False
|
||||
if result is not None and not isinstance(result, ModelResponseStream):
|
||||
return True
|
||||
return (
|
||||
"async_complete_streaming_response" in self.model_call_details
|
||||
or self.model_call_details.get("complete_streaming_response") is not None
|
||||
)
|
||||
|
||||
async def dispatch_success_handlers(
|
||||
self,
|
||||
result=None,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
cache_hit=None,
|
||||
prefer_async_handlers: bool = False,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""Route success logging to async and/or sync handlers for this request.
|
||||
|
||||
``prefer_async_handlers`` only bypasses the sync-SDK-only shortcut (e.g.
|
||||
``async for`` on a stream from ``completion()``). Legacy string callbacks
|
||||
still run via ``executor.submit(success_handler)`` when configured.
|
||||
"""
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
if self._is_assembled_stream_success(result):
|
||||
if self.model_call_details.get("has_dispatched_final_stream_success"):
|
||||
return
|
||||
self.model_call_details["has_dispatched_final_stream_success"] = True
|
||||
|
||||
litellm_params = self.model_call_details.get("litellm_params", {}) or {}
|
||||
sync_sdk = self._is_sync_litellm_request(litellm_params)
|
||||
passthrough = self.call_type == CallTypes.pass_through.value
|
||||
if sync_sdk and not prefer_async_handlers and not passthrough:
|
||||
self.success_handler(
|
||||
result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
await self.async_success_handler(
|
||||
result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
if not self._should_run_sync_callbacks_for_async_calls():
|
||||
return
|
||||
|
||||
executor.submit(
|
||||
self.success_handler,
|
||||
result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def should_run_logging(
|
||||
self,
|
||||
event_type: Literal[
|
||||
|
|
@ -2022,13 +2106,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
standard_logging_object=kwargs.get("standard_logging_object", None),
|
||||
)
|
||||
litellm_params = self.model_call_details.get("litellm_params", {})
|
||||
is_sync_request = (
|
||||
litellm_params.get(CallTypes.acompletion.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aresponses.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aembedding.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
|
||||
and litellm_params.get(CallTypes.atranscription.value, False) is not True
|
||||
)
|
||||
is_sync_request = self._is_sync_litellm_request(litellm_params)
|
||||
try:
|
||||
## BUILD COMPLETE STREAMED RESPONSE
|
||||
complete_streaming_response: Optional[
|
||||
|
|
@ -2484,9 +2562,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
print_verbose(
|
||||
"Logging Details LiteLLM-Async Success Call, cache_hit={}".format(cache_hit)
|
||||
)
|
||||
if not self.should_run_logging(
|
||||
if not self._is_assembled_stream_success(
|
||||
result
|
||||
) and not self.should_run_logging(
|
||||
event_type="async_success"
|
||||
): # prevent double logging
|
||||
): # prevent double logging (non-streaming)
|
||||
return
|
||||
|
||||
## CALCULATE COST FOR BATCH JOBS
|
||||
|
|
@ -2936,13 +3016,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
): # prevent double logging
|
||||
return
|
||||
litellm_params = self.model_call_details.get("litellm_params", {})
|
||||
is_sync_request = (
|
||||
litellm_params.get(CallTypes.acompletion.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aresponses.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aembedding.value, False) is not True
|
||||
and litellm_params.get(CallTypes.aimage_generation.value, False) is not True
|
||||
and litellm_params.get(CallTypes.atranscription.value, False) is not True
|
||||
)
|
||||
is_sync_request = self._is_sync_litellm_request(litellm_params)
|
||||
|
||||
try:
|
||||
start_time, end_time = self._failure_handler_helper_fn(
|
||||
|
|
|
|||
|
|
@ -1808,8 +1808,10 @@ class CustomStreamWrapper:
|
|||
processed_chunk, None, None, cache_hit
|
||||
)
|
||||
)
|
||||
## SYNC LOGGING
|
||||
self.logging_obj.success_handler(processed_chunk, None, None, cache_hit)
|
||||
## SYNC LOGGING — only for sync SDK entrypoints; async proxy paths export via async_success_handler
|
||||
litellm_params = self.logging_obj.model_call_details.get("litellm_params", {})
|
||||
if self.logging_obj._is_sync_litellm_request(litellm_params):
|
||||
self.logging_obj.success_handler(processed_chunk, None, None, cache_hit)
|
||||
|
||||
def finish_reason_handler(self):
|
||||
model_response = self.model_response_creator()
|
||||
|
|
@ -2206,23 +2208,19 @@ class CustomStreamWrapper:
|
|||
cache_hit,
|
||||
)
|
||||
else:
|
||||
# prefer_async_handlers routes CustomLogger to async_success_handler
|
||||
# when consumers use ``async for`` on sync-SDK streams. Legacy string
|
||||
# callbacks still run via executor.submit inside dispatch_success_handlers.
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(
|
||||
self.logging_obj.dispatch_success_handlers(
|
||||
complete_streaming_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
|
||||
executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
complete_streaming_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
)
|
||||
|
||||
raise StopAsyncIteration # Re-raise StopIteration
|
||||
else:
|
||||
self.sent_last_chunk = True
|
||||
|
|
|
|||
|
|
@ -2713,13 +2713,16 @@ class UserAPIKeyAuth(
|
|||
1. Regular API keys from LiteLLM DB
|
||||
2. JWT tokens used for connecting to LiteLLM API
|
||||
"""
|
||||
if api_key.startswith("sk-"):
|
||||
return hash_token(api_key)
|
||||
normalized = api_key
|
||||
if normalized[:7].lower() == "bearer ":
|
||||
normalized = normalized[7:]
|
||||
if normalized.startswith("sk-"):
|
||||
return hash_token(normalized)
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
|
||||
if JWTHandler.is_jwt(token=api_key):
|
||||
return f"hashed-jwt-{hash_token(token=api_key)}"
|
||||
return api_key
|
||||
if JWTHandler.is_jwt(token=normalized):
|
||||
return f"hashed-jwt-{hash_token(token=normalized)}"
|
||||
return normalized
|
||||
|
||||
@classmethod
|
||||
def get_litellm_internal_health_check_user_api_key_auth(cls) -> "UserAPIKeyAuth":
|
||||
|
|
|
|||
|
|
@ -1266,7 +1266,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# (ProxyLogging._fire_deferred_stream_logging) fires the
|
||||
# closure after the full streaming pipeline finishes.
|
||||
# The closure runs non-apply_guardrail hooks on the
|
||||
# assembled response, then fires both logging handlers.
|
||||
# assembled response, then fires success logging.
|
||||
# Only for CustomStreamWrapper — raw async generators from
|
||||
# passthrough routes bypass CSW and would orphan the closure.
|
||||
from litellm.litellm_core_utils.streaming_handler import (
|
||||
|
|
@ -1387,33 +1387,18 @@ class ProxyBaseLLMRequestProcessing:
|
|||
logging_obj._on_deferred_stream_complete = None # type: ignore[union-attr]
|
||||
try:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_success_handler(
|
||||
logging_obj.dispatch_success_handlers(
|
||||
response,
|
||||
cache_hit=None,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Error in orphaned streaming async logging: %s", e
|
||||
)
|
||||
try:
|
||||
from litellm.litellm_core_utils.thread_pool_executor import (
|
||||
executor as _exc,
|
||||
)
|
||||
|
||||
_exc.submit(
|
||||
logging_obj.success_handler,
|
||||
response,
|
||||
cache_hit=None,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Error in orphaned streaming sync logging: %s", e
|
||||
)
|
||||
|
||||
# Always return the client-requested model name (not provider-prefixed internal identifiers)
|
||||
# for OpenAI-compatible responses.
|
||||
|
|
@ -1615,7 +1600,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
) -> None:
|
||||
"""
|
||||
Run non-streaming post-call guardrail hooks on an assembled streaming
|
||||
response, then fire both async and sync logging handlers.
|
||||
response, then fire success logging via ``dispatch_success_handlers``.
|
||||
|
||||
Called by ProxyLogging._fire_deferred_stream_logging after the full
|
||||
streaming pipeline (including unified_guardrail end-of-stream blocks)
|
||||
|
|
@ -1631,8 +1616,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
Extracted as a static method so tests can call the production
|
||||
implementation directly rather than reimplementing the closure.
|
||||
"""
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
_response = assembled_response
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router as _global_llm_router
|
||||
|
|
@ -1691,31 +1674,23 @@ class ProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
finally:
|
||||
try:
|
||||
# Proxy streaming always runs in async context and proxy spend
|
||||
# logging is async-only; force async dispatch so DB/spend
|
||||
# callbacks fire regardless of the call-type heuristic in
|
||||
# _is_sync_litellm_request (which only recognizes a subset of
|
||||
# async markers stored in litellm_params).
|
||||
asyncio.create_task(
|
||||
captured_logging_obj.async_success_handler(
|
||||
captured_logging_obj.dispatch_success_handlers(
|
||||
_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Error in deferred streaming async logging: %s",
|
||||
e,
|
||||
)
|
||||
|
||||
try:
|
||||
executor.submit(
|
||||
captured_logging_obj.success_handler,
|
||||
_response,
|
||||
cache_hit=cache_hit,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Error in deferred streaming sync logging: %s",
|
||||
"Error in deferred streaming success logging: %s",
|
||||
e,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -414,6 +414,72 @@ class ResetBudgetJob:
|
|||
)
|
||||
return [LiteLLM_EndUserTable(**row.dict()) for row in rows]
|
||||
|
||||
async def _write_key_reset_updates(
|
||||
self, updated_keys: List[LiteLLM_VerificationToken]
|
||||
) -> None:
|
||||
"""
|
||||
Write per-row {spend, budget_reset_at} updates for keys.
|
||||
|
||||
Avoids the batched full-model update path, which trips
|
||||
prisma.errors.DataError on any row carrying object_permission_id or
|
||||
budget_limits (see #27730). Both fields are rejected by Prisma's
|
||||
update input type for LiteLLM_VerificationToken, and the failure
|
||||
aborts the entire batch — silently leaving spend over the cap and
|
||||
budget_reset_at unchanged forever.
|
||||
"""
|
||||
batcher = self.prisma_client.db.batch_()
|
||||
for k in updated_keys:
|
||||
token = getattr(k, "token", None)
|
||||
if token is None:
|
||||
continue
|
||||
batcher.litellm_verificationtoken.update(
|
||||
where={"token": token},
|
||||
data={"spend": 0, "budget_reset_at": k.budget_reset_at},
|
||||
)
|
||||
await batcher.commit()
|
||||
|
||||
async def _write_user_reset_updates(
|
||||
self, updated_users: List[LiteLLM_UserTable]
|
||||
) -> None:
|
||||
"""
|
||||
Write per-row {spend, budget_reset_at} updates for users.
|
||||
|
||||
Mirrors _write_key_reset_updates — avoids the full-model update path
|
||||
that trips Prisma's DataError on rows carrying unrecognised fields
|
||||
(see #27730).
|
||||
"""
|
||||
batcher = self.prisma_client.db.batch_()
|
||||
for u in updated_users:
|
||||
user_id = getattr(u, "user_id", None)
|
||||
if user_id is None:
|
||||
continue
|
||||
batcher.litellm_usertable.update(
|
||||
where={"user_id": user_id},
|
||||
data={"spend": 0, "budget_reset_at": u.budget_reset_at},
|
||||
)
|
||||
await batcher.commit()
|
||||
|
||||
async def _write_team_reset_updates(
|
||||
self, updated_teams: List[LiteLLM_TeamTable]
|
||||
) -> None:
|
||||
"""
|
||||
Write per-row {spend, budget_reset_at} updates for teams.
|
||||
|
||||
Mirrors _write_key_reset_updates — avoids the full-model update path
|
||||
that trips Prisma's DataError on rows carrying unrecognised fields
|
||||
(see #27730).
|
||||
"""
|
||||
batcher = self.prisma_client.db.batch_()
|
||||
for t in updated_teams:
|
||||
team_id = getattr(t, "team_id", None)
|
||||
if team_id is None:
|
||||
continue
|
||||
batcher.litellm_teamtable.update(
|
||||
where={"team_id": team_id},
|
||||
data={"spend": 0, "budget_reset_at": t.budget_reset_at},
|
||||
)
|
||||
await batcher.commit()
|
||||
|
||||
async def reset_budget_for_litellm_keys(self):
|
||||
"""
|
||||
Resets the budget for all the litellm keys
|
||||
|
|
@ -455,11 +521,7 @@ class ResetBudgetJob:
|
|||
)
|
||||
|
||||
if updated_keys:
|
||||
await self.prisma_client.update_data(
|
||||
query_type="update_many",
|
||||
data_list=updated_keys,
|
||||
table_name="key",
|
||||
)
|
||||
await self._write_key_reset_updates(updated_keys=updated_keys)
|
||||
for k in updated_keys:
|
||||
token = getattr(k, "token", None)
|
||||
if token:
|
||||
|
|
@ -544,11 +606,7 @@ class ResetBudgetJob:
|
|||
"Updated users %s", json.dumps(updated_users, indent=4, default=str)
|
||||
)
|
||||
if updated_users:
|
||||
await self.prisma_client.update_data(
|
||||
query_type="update_many",
|
||||
data_list=updated_users,
|
||||
table_name="user",
|
||||
)
|
||||
await self._write_user_reset_updates(updated_users=updated_users)
|
||||
for u in updated_users:
|
||||
user_id = getattr(u, "user_id", None)
|
||||
if user_id:
|
||||
|
|
@ -641,11 +699,7 @@ class ResetBudgetJob:
|
|||
"Updated teams %s", json.dumps(updated_teams, indent=4, default=str)
|
||||
)
|
||||
if updated_teams:
|
||||
await self.prisma_client.update_data(
|
||||
query_type="update_many",
|
||||
data_list=updated_teams,
|
||||
table_name="team",
|
||||
)
|
||||
await self._write_team_reset_updates(updated_teams=updated_teams)
|
||||
for t in updated_teams:
|
||||
team_id = getattr(t, "team_id", None)
|
||||
if team_id:
|
||||
|
|
@ -816,49 +870,16 @@ class ResetBudgetJob:
|
|||
"""
|
||||
In-place, updates spend=0, and sets budget_reset_at to current_time + budget_duration
|
||||
|
||||
Common logic for resetting budget for a team, user, or key
|
||||
Common logic for resetting budget for a team, user, or key.
|
||||
|
||||
Spend-counter invalidation happens in the caller, AFTER the DB write
|
||||
commits. Zeroing the counter here would open a bypass window when the
|
||||
DB write fails: get_current_spend reads 0 from Redis while the DB
|
||||
still holds the pre-reset value, admitting requests past the cap.
|
||||
"""
|
||||
try:
|
||||
item.spend = 0.0
|
||||
|
||||
# Reset the cross-pod spend counter.
|
||||
# Reset Redis directly (not via DualCache) so a Redis failure
|
||||
# doesn't silently leave a stale counter that get_current_spend
|
||||
# would read as authoritative, permanently blocking the user.
|
||||
from litellm.proxy.proxy_server import spend_counter_cache
|
||||
|
||||
counter_key = None
|
||||
if item_type == "key" and hasattr(item, "token") and item.token is not None: # type: ignore[union-attr]
|
||||
counter_key = f"spend:key:{item.token}" # type: ignore[union-attr]
|
||||
elif (
|
||||
item_type == "team"
|
||||
and hasattr(item, "team_id")
|
||||
and item.team_id is not None # type: ignore[union-attr]
|
||||
):
|
||||
counter_key = f"spend:team:{item.team_id}" # type: ignore[union-attr]
|
||||
|
||||
if counter_key is not None:
|
||||
# Always reset in-memory (local fallback)
|
||||
spend_counter_cache.in_memory_cache.set_cache(
|
||||
key=counter_key, value=0.0
|
||||
)
|
||||
# Explicitly reset Redis with warning on failure
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
await spend_counter_cache.redis_cache.async_set_cache(
|
||||
key=counter_key, value=0.0
|
||||
)
|
||||
except Exception as redis_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to reset spend counter in Redis for %s key=%s: %s. "
|
||||
"Budget may be over-enforced until counter expires.",
|
||||
item_type,
|
||||
counter_key,
|
||||
redis_err,
|
||||
)
|
||||
|
||||
if hasattr(item, "budget_duration") and item.budget_duration is not None:
|
||||
# Get standardized reset time based on budget duration
|
||||
from litellm.proxy.common_utils.timezone_utils import (
|
||||
get_budget_reset_time,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -224,6 +224,17 @@ TPM_RESERVED_SCOPES_KEY = "_litellm_tpm_reserved_scopes"
|
|||
# (e.g. async_log_failure_event firing after async_post_call_failure_hook)
|
||||
# does not double-refund.
|
||||
TPM_RESERVATION_RELEASED_KEY = "_litellm_tpm_reservation_released"
|
||||
RATE_LIMIT_DESCRIPTORS_KEY = "_litellm_rate_limit_descriptors"
|
||||
# Stash keys live ONLY in metadata channels — never at the top level of the
|
||||
# request body. Top-level keys are forwarded as body params to upstream
|
||||
# providers, which reject unknown fields with 400/429 errors.
|
||||
_LITELLM_STASH_KEYS: Tuple[str, ...] = (
|
||||
TPM_RESERVED_TOKENS_KEY,
|
||||
TPM_RESERVED_MODEL_KEY,
|
||||
TPM_RESERVED_SCOPES_KEY,
|
||||
TPM_RESERVATION_RELEASED_KEY,
|
||||
RATE_LIMIT_DESCRIPTORS_KEY,
|
||||
)
|
||||
|
||||
|
||||
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
|
||||
|
|
@ -1892,6 +1903,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"""
|
||||
verbose_proxy_logger.debug("Inside Rate Limit Pre-Call Hook")
|
||||
|
||||
# Reject caller-supplied stash values before any read/write. Otherwise
|
||||
# a client can inject ``_litellm_rate_limit_descriptors`` /
|
||||
# ``_litellm_tpm_reserved_tokens`` in body ``metadata`` and have
|
||||
# ``async_post_call_failure_hook`` refund TPM counters against scopes
|
||||
# they name (e.g. another tenant's api_key).
|
||||
self._strip_stash_keys_from_all_channels(data)
|
||||
|
||||
#########################################################
|
||||
# Check if the call type has a specific rate limiter
|
||||
# eg. for Batch APIs we need to use the batch rate limiter to read the input file and count the tokens and requests
|
||||
|
|
@ -2024,7 +2042,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
descriptors=descriptors,
|
||||
)
|
||||
else:
|
||||
data["_litellm_rate_limit_descriptors"] = descriptors
|
||||
self._stash_value_in_metadata_channels(
|
||||
data=data,
|
||||
key=RATE_LIMIT_DESCRIPTORS_KEY,
|
||||
value=descriptors,
|
||||
)
|
||||
# Capture the exact (key, value) scopes the reservation
|
||||
# incremented so post-call reconciliation only applies
|
||||
# the (actual - reserved) delta to those — unreserved
|
||||
|
|
@ -2059,6 +2081,29 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
f"TPM tokens reserved: {estimated_tokens} for model {requested_model}"
|
||||
)
|
||||
|
||||
# Defense-in-depth: scrub any stash key that escaped onto data
|
||||
# top-level (stale cache hit, router pass, test fixture) before the
|
||||
# body is forwarded to the provider.
|
||||
self._strip_stash_keys_from_top_level(data)
|
||||
|
||||
@staticmethod
|
||||
def _strip_stash_keys_from_top_level(data: Any) -> None:
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
for stash_key in _LITELLM_STASH_KEYS:
|
||||
data.pop(stash_key, None)
|
||||
|
||||
@classmethod
|
||||
def _strip_stash_keys_from_all_channels(cls, data: Any) -> None:
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
cls._strip_stash_keys_from_top_level(data)
|
||||
for channel in ("metadata", "litellm_metadata"):
|
||||
channel_dict = data.get(channel)
|
||||
if isinstance(channel_dict, dict):
|
||||
for stash_key in _LITELLM_STASH_KEYS:
|
||||
channel_dict.pop(stash_key, None)
|
||||
|
||||
def _create_pipeline_operations(
|
||||
self,
|
||||
key: str,
|
||||
|
|
@ -2233,18 +2278,29 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
return specified_rate_limit_type
|
||||
|
||||
@staticmethod
|
||||
def _stash_value_in_metadata_channels(
|
||||
data: Dict[str, Any],
|
||||
key: str,
|
||||
value: Any,
|
||||
) -> None:
|
||||
for channel in ("metadata", "litellm_metadata"):
|
||||
existing = data.get(channel)
|
||||
if isinstance(existing, dict):
|
||||
existing[key] = value
|
||||
elif channel == "metadata":
|
||||
# ``litellm_metadata`` is owned by the router; don't conjure
|
||||
# it here.
|
||||
data[channel] = {key: value}
|
||||
|
||||
@classmethod
|
||||
def _stash_reservation_in_data(
|
||||
cls,
|
||||
data: Dict[str, Any],
|
||||
estimated_tokens: int,
|
||||
reserved_model: Optional[str],
|
||||
reserved_scopes: Optional[List[Tuple[str, str]]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Persist the reservation amount, model, and reserved scopes into every
|
||||
channel a callback might read from: top-level kwargs (via ``**data``),
|
||||
request metadata, and litellm_metadata. Keeps reservation and
|
||||
reconciliation in sync.
|
||||
|
||||
``reserved_scopes`` is serialized as a list of [key, value] pairs so
|
||||
it round-trips through JSON-based metadata transports.
|
||||
"""
|
||||
|
|
@ -2252,30 +2308,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
[[k, v] for k, v in reserved_scopes] if reserved_scopes else None
|
||||
)
|
||||
|
||||
data[TPM_RESERVED_TOKENS_KEY] = estimated_tokens
|
||||
cls._stash_value_in_metadata_channels(
|
||||
data=data, key=TPM_RESERVED_TOKENS_KEY, value=estimated_tokens
|
||||
)
|
||||
if reserved_model:
|
||||
data[TPM_RESERVED_MODEL_KEY] = reserved_model
|
||||
cls._stash_value_in_metadata_channels(
|
||||
data=data, key=TPM_RESERVED_MODEL_KEY, value=reserved_model
|
||||
)
|
||||
if scopes_payload is not None:
|
||||
data[TPM_RESERVED_SCOPES_KEY] = scopes_payload
|
||||
|
||||
for channel in ("metadata", "litellm_metadata"):
|
||||
existing = data.get(channel)
|
||||
if isinstance(existing, dict):
|
||||
existing[TPM_RESERVED_TOKENS_KEY] = estimated_tokens
|
||||
if reserved_model:
|
||||
existing[TPM_RESERVED_MODEL_KEY] = reserved_model
|
||||
if scopes_payload is not None:
|
||||
existing[TPM_RESERVED_SCOPES_KEY] = scopes_payload
|
||||
elif channel == "metadata":
|
||||
# Only auto-create ``metadata`` (preserves prior behavior);
|
||||
# ``litellm_metadata`` is set by the router and shouldn't be
|
||||
# conjured here.
|
||||
stash: Dict[str, Any] = {TPM_RESERVED_TOKENS_KEY: estimated_tokens}
|
||||
if reserved_model:
|
||||
stash[TPM_RESERVED_MODEL_KEY] = reserved_model
|
||||
if scopes_payload is not None:
|
||||
stash[TPM_RESERVED_SCOPES_KEY] = scopes_payload
|
||||
data[channel] = stash
|
||||
cls._stash_value_in_metadata_channels(
|
||||
data=data, key=TPM_RESERVED_SCOPES_KEY, value=scopes_payload
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _lookup_stashed_value(
|
||||
|
|
@ -2284,19 +2327,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
key: str,
|
||||
) -> Any:
|
||||
"""
|
||||
Resolve a stashed value from any of the channels the request data can
|
||||
flow through to a callback.
|
||||
|
||||
Checks (in priority order):
|
||||
1. kwargs (top-level data fields propagate via **data)
|
||||
2. kwargs["litellm_params"]["metadata"] (request metadata channel)
|
||||
3. standard_logging_metadata (covers tests that mock the SLO directly)
|
||||
Resolve a stashed value from any metadata channel the request data
|
||||
can flow through to a callback. Top-level ``kwargs`` is not checked
|
||||
because stash keys must never live there.
|
||||
"""
|
||||
candidate = kwargs.get(key) if isinstance(kwargs, dict) else None
|
||||
if candidate is None:
|
||||
litellm_params = (
|
||||
kwargs.get("litellm_params") if isinstance(kwargs, dict) else None
|
||||
)
|
||||
candidate: Any = None
|
||||
if isinstance(kwargs, dict):
|
||||
for channel in ("metadata", "litellm_metadata"):
|
||||
channel_dict = kwargs.get(channel)
|
||||
if isinstance(channel_dict, dict) and key in channel_dict:
|
||||
candidate = channel_dict.get(key)
|
||||
if candidate is not None:
|
||||
return candidate
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
if isinstance(litellm_params, dict):
|
||||
lp_metadata = litellm_params.get("metadata")
|
||||
if isinstance(lp_metadata, dict):
|
||||
|
|
@ -2390,7 +2433,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"""
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
data[TPM_RESERVATION_RELEASED_KEY] = True
|
||||
for channel in ("metadata", "litellm_metadata"):
|
||||
existing = data.get(channel)
|
||||
if isinstance(existing, dict):
|
||||
|
|
@ -2811,9 +2853,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
return
|
||||
|
||||
# Refund directly against the descriptors we reserved against —
|
||||
# the pre-call hook stashes them on the request data before
|
||||
# success/failure callbacks run.
|
||||
stashed = request_data.get("_litellm_rate_limit_descriptors")
|
||||
# the pre-call hook stashes them in the request-data metadata
|
||||
# channels before success/failure callbacks run.
|
||||
stashed = self._lookup_stashed_value(
|
||||
kwargs=request_data,
|
||||
standard_logging_metadata=None,
|
||||
key=RATE_LIMIT_DESCRIPTORS_KEY,
|
||||
)
|
||||
descriptors: List[RateLimitDescriptor] = (
|
||||
stashed if isinstance(stashed, list) else []
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,13 +1,12 @@
|
|||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
from litellm.proxy._types import PassThroughEndpointLoggingResultValues
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
|
||||
|
|
@ -115,30 +114,81 @@ class PassThroughStreamingHandler:
|
|||
- OpenAI
|
||||
"""
|
||||
try:
|
||||
all_chunks = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(
|
||||
raw_bytes
|
||||
(
|
||||
standard_logging_response_object,
|
||||
kwargs,
|
||||
) = PassThroughStreamingHandler._build_passthrough_logging_result(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
request_body=request_body,
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
raw_bytes=raw_bytes,
|
||||
end_time=end_time,
|
||||
model=model,
|
||||
)
|
||||
standard_logging_response_object: Optional[
|
||||
PassThroughEndpointLoggingResultValues
|
||||
] = None
|
||||
kwargs: dict = {}
|
||||
if endpoint_type == EndpointType.ANTHROPIC:
|
||||
anthropic_passthrough_logging_handler_result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
request_body=request_body,
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
all_chunks=all_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
standard_logging_response_object = (
|
||||
anthropic_passthrough_logging_handler_result["result"]
|
||||
)
|
||||
kwargs = anthropic_passthrough_logging_handler_result["kwargs"]
|
||||
elif endpoint_type == EndpointType.VERTEX_AI:
|
||||
vertex_passthrough_logging_handler_result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks(
|
||||
# Always reached from an async context (anthropic_messages,
|
||||
# google_genai, and proxy pass-through stream tasks). prefer_async_handlers
|
||||
# keeps async-only loggers running even when call_type isn't pass_through
|
||||
# and litellm_params lacks an async flag (e.g. aanthropic_messages).
|
||||
await litellm_logging_obj.dispatch_success_handlers(
|
||||
result=standard_logging_response_object,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
prefer_async_handlers=True,
|
||||
**kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"Error in _route_streaming_logging_to_handler: {str(e)}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_passthrough_logging_result(
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
passthrough_success_handler_obj: PassThroughEndpointLogging,
|
||||
url_route: str,
|
||||
request_body: dict,
|
||||
endpoint_type: EndpointType,
|
||||
start_time: datetime,
|
||||
raw_bytes: List[bytes],
|
||||
end_time: datetime,
|
||||
model: Optional[str],
|
||||
) -> Tuple[PassThroughEndpointLoggingResultValues, dict]:
|
||||
"""
|
||||
Synchronous, CPU-bound reconstruction of the standard logging payload
|
||||
from collected raw SSE bytes. Extracted from
|
||||
_route_streaming_logging_to_handler so the per-endpoint dispatch can
|
||||
be unit-tested in isolation. Still invoked synchronously on the event
|
||||
loop; an off-loop dispatch is a future change, not part of this PR.
|
||||
"""
|
||||
all_chunks = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(
|
||||
raw_bytes
|
||||
)
|
||||
standard_logging_response_object: Optional[
|
||||
PassThroughEndpointLoggingResultValues
|
||||
] = None
|
||||
kwargs: dict = {}
|
||||
if endpoint_type == EndpointType.ANTHROPIC:
|
||||
anthropic_passthrough_logging_handler_result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
request_body=request_body,
|
||||
endpoint_type=endpoint_type,
|
||||
start_time=start_time,
|
||||
all_chunks=all_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
standard_logging_response_object = (
|
||||
anthropic_passthrough_logging_handler_result["result"]
|
||||
)
|
||||
kwargs = anthropic_passthrough_logging_handler_result["kwargs"]
|
||||
elif endpoint_type == EndpointType.VERTEX_AI:
|
||||
vertex_passthrough_logging_handler_result = (
|
||||
VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
|
|
@ -149,12 +199,14 @@ class PassThroughStreamingHandler:
|
|||
end_time=end_time,
|
||||
model=model,
|
||||
)
|
||||
standard_logging_response_object = (
|
||||
vertex_passthrough_logging_handler_result["result"]
|
||||
)
|
||||
kwargs = vertex_passthrough_logging_handler_result["kwargs"]
|
||||
elif endpoint_type == EndpointType.OPENAI:
|
||||
openai_passthrough_logging_handler_result = OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks(
|
||||
)
|
||||
standard_logging_response_object = (
|
||||
vertex_passthrough_logging_handler_result["result"]
|
||||
)
|
||||
kwargs = vertex_passthrough_logging_handler_result["kwargs"]
|
||||
elif endpoint_type == EndpointType.OPENAI:
|
||||
openai_passthrough_logging_handler_result = (
|
||||
OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
passthrough_success_handler_obj=passthrough_success_handler_obj,
|
||||
url_route=url_route,
|
||||
|
|
@ -164,40 +216,17 @@ class PassThroughStreamingHandler:
|
|||
all_chunks=all_chunks,
|
||||
end_time=end_time,
|
||||
)
|
||||
standard_logging_response_object = (
|
||||
openai_passthrough_logging_handler_result["result"]
|
||||
)
|
||||
kwargs = openai_passthrough_logging_handler_result["kwargs"]
|
||||
)
|
||||
standard_logging_response_object = (
|
||||
openai_passthrough_logging_handler_result["result"]
|
||||
)
|
||||
kwargs = openai_passthrough_logging_handler_result["kwargs"]
|
||||
|
||||
if standard_logging_response_object is None:
|
||||
standard_logging_response_object = StandardPassThroughResponseObject(
|
||||
response=f"cannot parse chunks to standard response object. Chunks={all_chunks}"
|
||||
)
|
||||
await litellm_logging_obj.async_success_handler(
|
||||
result=standard_logging_response_object,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
**kwargs,
|
||||
)
|
||||
if (
|
||||
litellm_logging_obj._should_run_sync_callbacks_for_async_calls()
|
||||
is False
|
||||
):
|
||||
return
|
||||
|
||||
executor.submit(
|
||||
litellm_logging_obj.success_handler,
|
||||
result=standard_logging_response_object,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
start_time=start_time,
|
||||
**kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
f"Error in _route_streaming_logging_to_handler: {str(e)}"
|
||||
if standard_logging_response_object is None:
|
||||
standard_logging_response_object = StandardPassThroughResponseObject(
|
||||
response=f"cannot parse chunks to standard response object. Chunks={all_chunks}"
|
||||
)
|
||||
return standard_logging_response_object, kwargs
|
||||
|
||||
@staticmethod
|
||||
def _extract_model_for_cost_injection(
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|||
PassthroughStandardLoggingPayload,
|
||||
)
|
||||
from litellm.types.utils import StandardPassThroughResponseObject
|
||||
from litellm.utils import executor as thread_pool_executor
|
||||
|
||||
from .llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
|
|
@ -94,19 +93,15 @@ class PassThroughEndpointLogging:
|
|||
cache_hit: bool,
|
||||
**kwargs,
|
||||
):
|
||||
"""Helper function to handle both sync and async logging operations"""
|
||||
# Submit to thread pool for sync logging
|
||||
thread_pool_executor.submit(
|
||||
logging_obj.success_handler,
|
||||
standard_logging_response_object,
|
||||
start_time,
|
||||
end_time,
|
||||
cache_hit,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Handle async logging
|
||||
await logging_obj.async_success_handler(
|
||||
"""Log pass-through success via the shared async dispatch path."""
|
||||
# Always reached from pass_through_async_success_handler, which runs in
|
||||
# an async context. call_type is "pass_through_endpoint" here, so the
|
||||
# passthrough guard in dispatch_success_handlers already forces the
|
||||
# async handler to run; pass prefer_async_handlers explicitly to match
|
||||
# the streaming sibling (_route_streaming_logging_to_handler) and keep
|
||||
# async-only loggers (e.g. the proxy spend logger) firing regardless of
|
||||
# how the call-type classification evolves.
|
||||
await logging_obj.dispatch_success_handlers(
|
||||
result=(
|
||||
json.dumps(result)
|
||||
if isinstance(result, dict)
|
||||
|
|
@ -115,6 +110,7 @@ class PassThroughEndpointLogging:
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
prefer_async_handlers=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -3148,6 +3148,7 @@ all_litellm_params = (
|
|||
"allowed_openai_params",
|
||||
"litellm_session_id",
|
||||
"use_litellm_proxy",
|
||||
"use_chat_completions_api",
|
||||
"prompt_label",
|
||||
"shared_session",
|
||||
"search_tool_name",
|
||||
|
|
|
|||
|
|
@ -22,6 +22,60 @@ from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
|
|||
# In a real-world scenario, these would be instances of LiteLLM_VerificationToken, LiteLLM_UserTable, etc.
|
||||
|
||||
|
||||
def _attrify(d: dict):
|
||||
"""
|
||||
Wrap a dict so that attribute access (`.token`, `.user_id`, `.team_id`,
|
||||
etc.) works alongside the existing item-access the fake_reset_* helpers
|
||||
rely on. The reset job's narrow-write helpers use `getattr(item, "token",
|
||||
None)` (et al), which returns None for plain dicts — that would silently
|
||||
skip the row.
|
||||
"""
|
||||
class _AttrDict(dict):
|
||||
def __getattr__(self, k):
|
||||
try:
|
||||
return self[k]
|
||||
except KeyError:
|
||||
raise AttributeError(k)
|
||||
|
||||
def __setattr__(self, k, v):
|
||||
self[k] = v
|
||||
|
||||
return _AttrDict(d)
|
||||
|
||||
|
||||
def _wire_batcher_for_test(prisma_client):
|
||||
"""
|
||||
Wire prisma_client.db.batch_() to return a mock batcher whose .commit() is
|
||||
awaitable and whose per-table .update() calls get captured. The reset job
|
||||
writes key/user/team resets via prisma.db.batch_().<table>.update — not via
|
||||
prisma_client.update_data — so tests must let that batch path complete.
|
||||
|
||||
Returns the list that will accumulate {table, where, data} dicts from
|
||||
each captured update call.
|
||||
"""
|
||||
batch_calls = []
|
||||
|
||||
def make_batcher():
|
||||
class _Table:
|
||||
def __init__(self, table_name):
|
||||
self._table_name = table_name
|
||||
|
||||
def update(self, where=None, data=None):
|
||||
batch_calls.append(
|
||||
{"table": self._table_name, "where": where, "data": data}
|
||||
)
|
||||
|
||||
batcher = MagicMock()
|
||||
batcher.litellm_verificationtoken = _Table("key")
|
||||
batcher.litellm_usertable = _Table("user")
|
||||
batcher.litellm_teamtable = _Table("team")
|
||||
batcher.commit = AsyncMock(return_value=None)
|
||||
return batcher
|
||||
|
||||
prisma_client.db.batch_ = MagicMock(side_effect=make_batcher)
|
||||
return batch_calls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reset_budget_keys_partial_failure():
|
||||
"""
|
||||
|
|
@ -45,6 +99,9 @@ async def test_reset_budget_keys_partial_failure():
|
|||
return_value=[key1, key2, key3, key4, key5, key6]
|
||||
)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
# Reset job writes key resets via prisma.db.batch_().<table>.update — not
|
||||
# via update_data — so wire that path.
|
||||
batch_calls = _wire_batcher_for_test(prisma_client)
|
||||
|
||||
# Using a dummy logging object with async hooks mocked out.
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
|
@ -56,6 +113,15 @@ async def test_reset_budget_keys_partial_failure():
|
|||
|
||||
now = datetime.utcnow()
|
||||
|
||||
# token is needed because the new write path uses where={"token": ...}
|
||||
# and _AttrDict makes getattr work alongside item access used by fake_reset_key.
|
||||
for k in [key1, key2, key3, key4, key5, key6]:
|
||||
k.setdefault("token", k["id"])
|
||||
key1, key2, key3, key4, key5, key6 = (
|
||||
_attrify(k) for k in [key1, key2, key3, key4, key5, key6]
|
||||
)
|
||||
prisma_client.get_data = AsyncMock(return_value=[key1, key2, key3, key4, key5, key6])
|
||||
|
||||
async def fake_reset_key(key, current_time):
|
||||
if key["id"] == "key1":
|
||||
# Simulate a failure on key1 (for example, this might be due to an invariant check)
|
||||
|
|
@ -80,17 +146,17 @@ async def test_reset_budget_keys_partial_failure():
|
|||
# Assert that the helper was called for 6 keys
|
||||
assert mock_reset_key.call_count == 6
|
||||
|
||||
# Assert that update_data was called once with a list containing all 6 keys
|
||||
prisma_client.update_data.assert_awaited_once()
|
||||
update_call = prisma_client.update_data.call_args
|
||||
assert update_call.kwargs.get("table_name") == "key"
|
||||
updated_keys = update_call.kwargs.get("data_list", [])
|
||||
assert len(updated_keys) == 5
|
||||
assert updated_keys[0]["id"] == "key2"
|
||||
assert updated_keys[1]["id"] == "key3"
|
||||
assert updated_keys[2]["id"] == "key4"
|
||||
assert updated_keys[3]["id"] == "key5"
|
||||
assert updated_keys[4]["id"] == "key6"
|
||||
# Assert that the new narrow write path got 5 batched updates (key1 failed).
|
||||
# update_data must NOT have been called for keys.
|
||||
prisma_client.update_data.assert_not_awaited()
|
||||
key_writes = [c for c in batch_calls if c["table"] == "key"]
|
||||
assert len(key_writes) == 5
|
||||
written_ids = [c["where"]["token"] for c in key_writes]
|
||||
assert written_ids == ["key2", "key3", "key4", "key5", "key6"]
|
||||
# And every write must carry only {spend, budget_reset_at} — never the full row.
|
||||
for c in key_writes:
|
||||
assert set(c["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
assert c["data"]["spend"] == 0
|
||||
|
||||
# Verify that the failure logging hook was scheduled (due to the failure for key1)
|
||||
failure_hook_calls = (
|
||||
|
|
@ -125,6 +191,7 @@ async def test_reset_budget_users_partial_failure():
|
|||
return_value=[user1, user2, user3, user4, user5, user6]
|
||||
)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
batch_calls = _wire_batcher_for_test(prisma_client)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -133,6 +200,15 @@ async def test_reset_budget_users_partial_failure():
|
|||
|
||||
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
|
||||
|
||||
# user_id required for the new write path's where clause; _AttrDict so
|
||||
# getattr(u, 'user_id') works alongside the dict access fake_reset_user uses.
|
||||
for u in [user1, user2, user3, user4, user5, user6]:
|
||||
u.setdefault("user_id", u["id"])
|
||||
user1, user2, user3, user4, user5, user6 = (
|
||||
_attrify(u) for u in [user1, user2, user3, user4, user5, user6]
|
||||
)
|
||||
prisma_client.get_data = AsyncMock(return_value=[user1, user2, user3, user4, user5, user6])
|
||||
|
||||
async def fake_reset_user(user, current_time):
|
||||
if user["id"] == "user1":
|
||||
raise Exception("Simulated failure for user1")
|
||||
|
|
@ -150,16 +226,14 @@ async def test_reset_budget_users_partial_failure():
|
|||
await asyncio.sleep(0.1)
|
||||
|
||||
assert mock_reset_user.call_count == 6
|
||||
prisma_client.update_data.assert_awaited_once()
|
||||
update_call = prisma_client.update_data.call_args
|
||||
assert update_call.kwargs.get("table_name") == "user"
|
||||
updated_users = update_call.kwargs.get("data_list", [])
|
||||
assert len(updated_users) == 5
|
||||
assert updated_users[0]["id"] == "user2"
|
||||
assert updated_users[1]["id"] == "user3"
|
||||
assert updated_users[2]["id"] == "user4"
|
||||
assert updated_users[3]["id"] == "user5"
|
||||
assert updated_users[4]["id"] == "user6"
|
||||
prisma_client.update_data.assert_not_awaited()
|
||||
user_writes = [c for c in batch_calls if c["table"] == "user"]
|
||||
assert len(user_writes) == 5
|
||||
written_ids = [c["where"]["user_id"] for c in user_writes]
|
||||
assert written_ids == ["user2", "user3", "user4", "user5", "user6"]
|
||||
for c in user_writes:
|
||||
assert set(c["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
assert c["data"]["spend"] == 0
|
||||
|
||||
failure_hook_calls = (
|
||||
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
|
||||
|
|
@ -308,6 +382,7 @@ async def test_reset_budget_teams_partial_failure():
|
|||
prisma_client = MagicMock()
|
||||
prisma_client.get_data = AsyncMock(return_value=[team1, team2])
|
||||
prisma_client.update_data = AsyncMock()
|
||||
batch_calls = _wire_batcher_for_test(prisma_client)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -316,6 +391,12 @@ async def test_reset_budget_teams_partial_failure():
|
|||
|
||||
job = ResetBudgetJob(proxy_logging_obj, prisma_client)
|
||||
|
||||
# team_id required for the new write path's where clause; _AttrDict for getattr.
|
||||
for t in [team1, team2]:
|
||||
t.setdefault("team_id", t["id"])
|
||||
team1, team2 = _attrify(team1), _attrify(team2)
|
||||
prisma_client.get_data = AsyncMock(return_value=[team1, team2])
|
||||
|
||||
async def fake_reset_team(team, current_time):
|
||||
if team["id"] == "team1":
|
||||
raise Exception("Simulated failure for team1")
|
||||
|
|
@ -333,12 +414,12 @@ async def test_reset_budget_teams_partial_failure():
|
|||
await asyncio.sleep(0.1)
|
||||
|
||||
assert mock_reset_team.call_count == 2
|
||||
prisma_client.update_data.assert_awaited_once()
|
||||
update_call = prisma_client.update_data.call_args
|
||||
assert update_call.kwargs.get("table_name") == "team"
|
||||
updated_teams = update_call.kwargs.get("data_list", [])
|
||||
assert len(updated_teams) == 1
|
||||
assert updated_teams[0]["id"] == "team2"
|
||||
prisma_client.update_data.assert_not_awaited()
|
||||
team_writes = [c for c in batch_calls if c["table"] == "team"]
|
||||
assert len(team_writes) == 1
|
||||
assert team_writes[0]["where"] == {"team_id": "team2"}
|
||||
assert set(team_writes[0]["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
assert team_writes[0]["data"]["spend"] == 0
|
||||
|
||||
failure_hook_calls = (
|
||||
proxy_logging_obj.service_logging_obj.async_service_failure_hook.call_args_list
|
||||
|
|
@ -402,6 +483,18 @@ async def test_reset_budget_continues_other_categories_on_failure():
|
|||
|
||||
prisma_client.get_data = AsyncMock(side_effect=fake_get_data)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
batch_calls = _wire_batcher_for_test(prisma_client)
|
||||
# ID fields required by the new write path's where clauses; _AttrDict
|
||||
# lets getattr() see them alongside the item-access fake_reset_* helpers use.
|
||||
for k in [key1, key2]:
|
||||
k.setdefault("token", k["id"])
|
||||
for u in [user1, user2]:
|
||||
u.setdefault("user_id", u["id"])
|
||||
for t in [team1, team2]:
|
||||
t.setdefault("team_id", t["id"])
|
||||
key1, key2 = _attrify(key1), _attrify(key2)
|
||||
user1, user2 = _attrify(user1), _attrify(user2)
|
||||
team1, team2 = _attrify(team1), _attrify(team2)
|
||||
# Mock db.litellm_verificationtoken.update_many (used by reset_budget_for_keys_linked_to_budgets)
|
||||
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
|
||||
return_value={"count": 0}
|
||||
|
|
@ -488,32 +581,29 @@ async def test_reset_budget_continues_other_categories_on_failure():
|
|||
"team_membership",
|
||||
}
|
||||
|
||||
# Verify that update_data was called three times (one per category, enduser update includes two)
|
||||
assert prisma_client.update_data.await_count == 5
|
||||
# After the fix, keys/users/teams write via prisma.db.batch_().<table>.update,
|
||||
# so only budget + enduser still go through update_data.
|
||||
calls = prisma_client.update_data.await_args_list
|
||||
|
||||
# Check keys update: both keys succeed.
|
||||
keys_call = calls[0]
|
||||
assert keys_call.kwargs.get("table_name") == "key"
|
||||
assert len(keys_call.kwargs.get("data_list", [])) == 2
|
||||
|
||||
# Check users update: only user2 succeeded.
|
||||
users_call = calls[1]
|
||||
assert users_call.kwargs.get("table_name") == "user"
|
||||
users_updated = users_call.kwargs.get("data_list", [])
|
||||
assert len(users_updated) == 1
|
||||
assert users_updated[0]["id"] == "user2"
|
||||
|
||||
# Check teams update: both teams succeed.
|
||||
teams_call = calls[2]
|
||||
assert teams_call.kwargs.get("table_name") == "team"
|
||||
assert len(teams_call.kwargs.get("data_list", [])) == 2
|
||||
update_data_tables = [c.kwargs.get("table_name") for c in calls]
|
||||
assert sorted(update_data_tables) == ["budget", "enduser"]
|
||||
|
||||
# Check enduser update: enduser succeed.
|
||||
enduser_call = calls[4]
|
||||
assert enduser_call.kwargs.get("table_name") == "enduser"
|
||||
enduser_call = next(c for c in calls if c.kwargs.get("table_name") == "enduser")
|
||||
assert len(enduser_call.kwargs.get("data_list", [])) == 1
|
||||
|
||||
# Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams.
|
||||
key_writes = [c for c in batch_calls if c["table"] == "key"]
|
||||
user_writes = [c for c in batch_calls if c["table"] == "user"]
|
||||
team_writes = [c for c in batch_calls if c["table"] == "team"]
|
||||
assert len(key_writes) == 2
|
||||
assert len(user_writes) == 1
|
||||
assert user_writes[0]["where"] == {"user_id": "user2"}
|
||||
assert len(team_writes) == 2
|
||||
# Every batched write must carry only the two reset fields, never the full row.
|
||||
for c in key_writes + user_writes + team_writes:
|
||||
assert set(c["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
assert c["data"]["spend"] == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Additional tests for service logger behavior (keys, users, teams, endusers)
|
||||
|
|
@ -527,12 +617,13 @@ async def test_service_logger_keys_success():
|
|||
logger success hook is called with the correct event metadata and no exception is logged.
|
||||
"""
|
||||
keys = [
|
||||
{"id": "key1", "spend": 10.0, "budget_duration": 60},
|
||||
{"id": "key2", "spend": 15.0, "budget_duration": 60},
|
||||
{"id": "key1", "spend": 10.0, "budget_duration": 60, "token": "key1"},
|
||||
{"id": "key2", "spend": 15.0, "budget_duration": 60, "token": "key2"},
|
||||
]
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.get_data = AsyncMock(return_value=keys)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
_wire_batcher_for_test(prisma_client)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -644,12 +735,13 @@ async def test_service_logger_users_success():
|
|||
the correct metadata and no exception is logged.
|
||||
"""
|
||||
users = [
|
||||
{"id": "user1", "spend": 20.0, "budget_duration": 120},
|
||||
{"id": "user2", "spend": 25.0, "budget_duration": 120},
|
||||
{"id": "user1", "spend": 20.0, "budget_duration": 120, "user_id": "user1"},
|
||||
{"id": "user2", "spend": 25.0, "budget_duration": 120, "user_id": "user2"},
|
||||
]
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.get_data = AsyncMock(return_value=users)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
_wire_batcher_for_test(prisma_client)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
@ -756,12 +848,13 @@ async def test_service_logger_teams_success():
|
|||
the proper metadata and nothing is logged as an exception.
|
||||
"""
|
||||
teams = [
|
||||
{"id": "team1", "spend": 30.0, "budget_duration": 180},
|
||||
{"id": "team2", "spend": 35.0, "budget_duration": 180},
|
||||
{"id": "team1", "spend": 30.0, "budget_duration": 180, "team_id": "team1"},
|
||||
{"id": "team2", "spend": 35.0, "budget_duration": 180, "team_id": "team2"},
|
||||
]
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.get_data = AsyncMock(return_value=teams)
|
||||
prisma_client.update_data = AsyncMock()
|
||||
_wire_batcher_for_test(prisma_client)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.service_logging_obj = MagicMock()
|
||||
|
|
|
|||
|
|
@ -97,6 +97,123 @@ async def test_chunk_processor_yields_raw_bytes(endpoint_type, url_route):
|
|||
), "Collected chunks do not match raw chunks"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_streaming_logging_runs_async_handler_for_sdk_passthrough():
|
||||
"""
|
||||
SDK pass-through streaming (anthropic_messages, google generate_content) must run
|
||||
the async success handler so async-only loggers record the assembled stream.
|
||||
|
||||
Regression for duplicate-trace dedupe: dispatch_success_handlers treated these as
|
||||
sync SDK requests because call_type is not ``pass_through_endpoint`` and
|
||||
litellm_params carries no ``acompletion`` flag, so only the sync success_handler
|
||||
ran and CustomLogger.async_log_success_event never fired.
|
||||
"""
|
||||
import time
|
||||
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type=CallTypes.anthropic_messages.value,
|
||||
start_time=time.time(),
|
||||
litellm_call_id="test-id",
|
||||
function_id="fn",
|
||||
)
|
||||
logging_obj.model_call_details["litellm_params"] = {"anthropic_messages": True}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
PassThroughStreamingHandler,
|
||||
"_build_passthrough_logging_result",
|
||||
return_value=({"id": "slp"}, {}),
|
||||
),
|
||||
patch.object(
|
||||
logging_obj, "async_success_handler", new_callable=AsyncMock
|
||||
) as mock_async,
|
||||
patch.object(
|
||||
logging_obj, "success_handler", new_callable=MagicMock
|
||||
) as mock_sync,
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_callbacks_for_async_calls",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
await PassThroughStreamingHandler._route_streaming_logging_to_handler(
|
||||
litellm_logging_obj=logging_obj,
|
||||
passthrough_success_handler_obj=MagicMock(),
|
||||
url_route="/v1/messages",
|
||||
request_body={},
|
||||
endpoint_type=EndpointType.ANTHROPIC,
|
||||
start_time=datetime.now(),
|
||||
raw_bytes=[],
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
mock_async.assert_awaited_once()
|
||||
mock_sync.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_logging_runs_async_handler_for_passthrough():
|
||||
"""
|
||||
Non-streaming pass-through logging (_handle_logging) must always run the
|
||||
async success handler so async-only loggers (e.g. the proxy spend logger)
|
||||
record the request.
|
||||
|
||||
_handle_logging is only ever reached from pass_through_async_success_handler
|
||||
(an async context), so it forces async dispatch via prefer_async_handlers.
|
||||
This pins that contract independent of the call-type classification: even a
|
||||
call_type that _is_sync_litellm_request would classify as sync (here
|
||||
"completion" with no async marker in litellm_params) must still reach
|
||||
async_success_handler. Without prefer_async_handlers=True the sync-only
|
||||
branch would return early and async_log_success_event would never fire.
|
||||
"""
|
||||
import time
|
||||
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type=CallTypes.completion.value,
|
||||
start_time=time.time(),
|
||||
litellm_call_id="test-id",
|
||||
function_id="fn",
|
||||
)
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
|
||||
handler = PassThroughEndpointLogging()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
logging_obj, "async_success_handler", new_callable=AsyncMock
|
||||
) as mock_async,
|
||||
patch.object(
|
||||
logging_obj, "success_handler", new_callable=MagicMock
|
||||
) as mock_sync,
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_callbacks_for_async_calls",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
await handler._handle_logging(
|
||||
logging_obj=logging_obj,
|
||||
standard_logging_response_object={"id": "slp"},
|
||||
result="",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
)
|
||||
|
||||
mock_async.assert_awaited_once()
|
||||
mock_sync.assert_not_called()
|
||||
|
||||
|
||||
def test_convert_raw_bytes_to_str_lines():
|
||||
"""
|
||||
Test that the _convert_raw_bytes_to_str_lines method correctly converts raw bytes to a list of strings
|
||||
|
|
|
|||
|
|
@ -95,6 +95,21 @@ router = Router(
|
|||
)
|
||||
|
||||
|
||||
def _register_proxy_test_logger(callback_logger: testLogger) -> None:
|
||||
"""
|
||||
Register the test logger on global callback lists.
|
||||
|
||||
``function_setup`` dedupes by object identity; each parametrized case
|
||||
constructs a new ``testLogger`` and must replace the global lists, not
|
||||
only ``litellm.callbacks``.
|
||||
"""
|
||||
litellm.callbacks = [callback_logger]
|
||||
litellm.success_callback = [callback_logger]
|
||||
litellm.failure_callback = [callback_logger]
|
||||
litellm._async_success_callback = [callback_logger]
|
||||
litellm._async_failure_callback = [callback_logger]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route, body",
|
||||
[
|
||||
|
|
@ -115,7 +130,7 @@ router = Router(
|
|||
"/v1/embeddings",
|
||||
{
|
||||
"input": "The food was delicious and the waiter...",
|
||||
"model": "text-embedding-ada-002",
|
||||
"model": "fake-model",
|
||||
"encoding_format": "float",
|
||||
},
|
||||
),
|
||||
|
|
@ -133,7 +148,7 @@ async def test_chat_completion_request_with_redaction(route, body):
|
|||
|
||||
setattr(proxy_server, "llm_router", router)
|
||||
_test_logger = testLogger()
|
||||
litellm.callbacks = [_test_logger]
|
||||
_register_proxy_test_logger(_test_logger)
|
||||
litellm.set_verbose = True
|
||||
|
||||
# Prepare the query string
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -786,6 +787,211 @@ def test_success_handler_runs_sync_callbacks_for_sync_requests(logging_obj, call
|
|||
dummy_logger.log_stream_event.assert_not_called()
|
||||
|
||||
|
||||
def test_is_sync_litellm_request():
|
||||
assert LitellmLogging._is_sync_litellm_request({}) is True
|
||||
assert LitellmLogging._is_sync_litellm_request({"acompletion": True}) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_success_handlers_invokes_callbacks_once_for_final_stream(
|
||||
logging_obj,
|
||||
):
|
||||
"""Second final-stream dispatch must not re-export (CSW + deferred guardrail paths)."""
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class MockCallback(CustomLogger):
|
||||
pass
|
||||
|
||||
mock_callback = MockCallback()
|
||||
original_async_callbacks = list(litellm._async_success_callback or [])
|
||||
litellm._async_success_callback = [mock_callback]
|
||||
|
||||
result = ModelResponse(
|
||||
id="resp-dedupe",
|
||||
model="gpt-4o-mini",
|
||||
choices=[
|
||||
{
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
)
|
||||
|
||||
try:
|
||||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["litellm_params"] = {"acompletion": True}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
mock_callback, "async_log_success_event", new_callable=AsyncMock
|
||||
) as mock_async_log,
|
||||
patch.object(mock_callback, "log_success_event") as mock_sync_log,
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_success_handler_helper_fn",
|
||||
return_value=(time.time(), time.time(), result),
|
||||
),
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_get_assembled_streaming_response",
|
||||
return_value=result,
|
||||
),
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_callbacks_for_async_calls",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
await logging_obj.dispatch_success_handlers(result=result)
|
||||
await logging_obj.dispatch_success_handlers(result=result)
|
||||
|
||||
mock_async_log.assert_awaited_once()
|
||||
mock_sync_log.assert_not_called()
|
||||
finally:
|
||||
litellm._async_success_callback = original_async_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_success_handlers_sync_path_invokes_callback_once_for_final_stream(
|
||||
logging_obj,
|
||||
):
|
||||
"""Sync dispatch path must also dedupe when dispatch is called twice."""
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class MockCallback(CustomLogger):
|
||||
pass
|
||||
|
||||
mock_callback = MockCallback()
|
||||
original_success_callbacks = list(litellm.success_callback or [])
|
||||
litellm.success_callback = [mock_callback]
|
||||
|
||||
result = ModelResponse(
|
||||
id="resp-sync-dedupe",
|
||||
model="gpt-4o-mini",
|
||||
choices=[
|
||||
{
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
)
|
||||
|
||||
try:
|
||||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
|
||||
with (
|
||||
patch.object(mock_callback, "log_success_event") as mock_sync_log,
|
||||
patch.object(
|
||||
mock_callback, "async_log_success_event", new_callable=AsyncMock
|
||||
) as mock_async_log,
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_success_handler_helper_fn",
|
||||
return_value=(time.time(), time.time(), result),
|
||||
),
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_get_assembled_streaming_response",
|
||||
return_value=result,
|
||||
),
|
||||
):
|
||||
await logging_obj.dispatch_success_handlers(result=result)
|
||||
await logging_obj.dispatch_success_handlers(result=result)
|
||||
|
||||
mock_sync_log.assert_called_once()
|
||||
mock_async_log.assert_not_awaited()
|
||||
finally:
|
||||
litellm.success_callback = original_success_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_prefer_async_handlers_runs_legacy_callbacks(
|
||||
logging_obj,
|
||||
):
|
||||
"""``prefer_async_handlers`` must not skip executor.submit for string callbacks."""
|
||||
result = ModelResponse(
|
||||
id="resp-prefer-async",
|
||||
model="gpt-4o-mini",
|
||||
choices=[
|
||||
{
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
logging_obj.stream = True
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
logging_obj, "async_success_handler", new_callable=AsyncMock
|
||||
) as mock_async,
|
||||
patch.object(
|
||||
logging_obj, "success_handler", new_callable=MagicMock
|
||||
) as mock_sync,
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_callbacks_for_async_calls",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.litellm_logging.executor.submit"
|
||||
) as mock_submit,
|
||||
):
|
||||
await logging_obj.dispatch_success_handlers(
|
||||
result=result,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
|
||||
mock_async.assert_awaited_once()
|
||||
mock_sync.assert_not_called()
|
||||
mock_submit.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_success_handlers_invokes_async_callback_for_pass_through(
|
||||
logging_obj,
|
||||
):
|
||||
"""Pass-through must use async_success_handler (CustomLogger skips sync success_handler)."""
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
class MockCallback(CustomLogger):
|
||||
pass
|
||||
|
||||
mock_callback = MockCallback()
|
||||
original_async_callbacks = list(litellm._async_success_callback or [])
|
||||
litellm._async_success_callback = [mock_callback]
|
||||
|
||||
logging_obj.call_type = CallTypes.pass_through.value
|
||||
logging_obj.stream = False
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(
|
||||
mock_callback, "async_log_success_event", new_callable=AsyncMock
|
||||
) as mock_async_log,
|
||||
patch.object(mock_callback, "log_success_event") as mock_sync_log,
|
||||
):
|
||||
await logging_obj.dispatch_success_handlers(result={"id": "pt-1"})
|
||||
|
||||
mock_async_log.assert_awaited_once()
|
||||
mock_sync_log.assert_not_called()
|
||||
finally:
|
||||
litellm._async_success_callback = original_async_callbacks
|
||||
|
||||
|
||||
def test_success_handler_skips_guardrail_logging_hook_when_disabled(logging_obj):
|
||||
"""Ensure CustomGuardrail logging_hook is skipped when should_run_guardrail is False."""
|
||||
import datetime
|
||||
|
|
@ -1351,7 +1557,7 @@ async def test_e2e_generate_cold_storage_object_key_with_custom_logger_s3_path()
|
|||
Test that _generate_cold_storage_object_key uses s3_path from custom logger instance.
|
||||
"""
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
|
||||
|
|
@ -1404,7 +1610,7 @@ async def test_e2e_generate_cold_storage_object_key_with_logger_no_s3_path():
|
|||
Test that _generate_cold_storage_object_key falls back to empty s3_path when logger has no s3_path.
|
||||
"""
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
|
||||
|
|
|
|||
|
|
@ -569,8 +569,6 @@ async def test_streaming_with_usage_and_logging(sync_mode: bool):
|
|||
== final_usage_block
|
||||
)
|
||||
|
||||
print(mock_log_success_event.call_args.kwargs.keys())
|
||||
|
||||
|
||||
def test_streaming_handler_with_stop_chunk(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper,
|
||||
|
|
@ -2036,23 +2034,19 @@ async def test_azure_streaming_role_preserved_with_include_usage(sync_mode: bool
|
|||
chunks.append(chunk)
|
||||
|
||||
# The prompt_filter chunk should be forwarded with choices=[]
|
||||
assert len(chunks[0].choices) == 0, (
|
||||
f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices"
|
||||
)
|
||||
assert (
|
||||
len(chunks[0].choices) == 0
|
||||
), f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices"
|
||||
|
||||
# At least one chunk must have role='assistant' in its delta
|
||||
has_role = any(
|
||||
len(c.choices) > 0
|
||||
and getattr(c.choices[0].delta, "role", None) == "assistant"
|
||||
len(c.choices) > 0 and getattr(c.choices[0].delta, "role", None) == "assistant"
|
||||
for c in chunks
|
||||
)
|
||||
assert has_role, (
|
||||
"No chunk contained role='assistant' in delta (issue #24221). "
|
||||
"Chunk deltas: "
|
||||
+ str([
|
||||
c.choices[0].delta if c.choices else "no choices"
|
||||
for c in chunks
|
||||
])
|
||||
+ str([c.choices[0].delta if c.choices else "no choices" for c in chunks])
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,74 @@
|
|||
"""
|
||||
Regression test for issue #28146.
|
||||
|
||||
`use_chat_completions_api` is a LiteLLM-internal control flag (it forces the
|
||||
/responses -> /chat/completions bridge). When set as a model-level param in the
|
||||
proxy config, it must never be forwarded to the upstream provider's request
|
||||
body. OpenAI/Anthropic reject unknown body params with HTTP 400.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../.."))
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import all_litellm_params
|
||||
from litellm.utils import get_non_default_completion_params
|
||||
|
||||
|
||||
def test_use_chat_completions_api_is_a_known_litellm_param():
|
||||
assert "use_chat_completions_api" in all_litellm_params
|
||||
|
||||
|
||||
def test_use_chat_completions_api_not_forwarded_as_provider_param():
|
||||
forwarded = get_non_default_completion_params(
|
||||
{"use_chat_completions_api": True, "temperature": 0.5}
|
||||
)
|
||||
assert "use_chat_completions_api" not in forwarded
|
||||
|
||||
|
||||
def test_completion_does_not_leak_flag_into_provider_request_body():
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump.return_value = {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion",
|
||||
"created": 1234567890,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 1,
|
||||
"total_tokens": 2,
|
||||
},
|
||||
}
|
||||
|
||||
mock_raw_response = MagicMock()
|
||||
mock_raw_response.headers = {}
|
||||
mock_raw_response.parse.return_value = mock_response
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.chat.completions.with_raw_response.create.return_value = (
|
||||
mock_raw_response
|
||||
)
|
||||
|
||||
litellm.completion(
|
||||
model="openai/gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
use_chat_completions_api=True,
|
||||
api_key="sk-test",
|
||||
client=mock_client,
|
||||
)
|
||||
|
||||
create_kwargs = (
|
||||
mock_client.chat.completions.with_raw_response.create.call_args.kwargs
|
||||
)
|
||||
assert "use_chat_completions_api" not in create_kwargs
|
||||
assert "use_chat_completions_api" not in (create_kwargs.get("extra_body") or {})
|
||||
|
|
@ -213,7 +213,7 @@ class TestMCPRequestHandler:
|
|||
# Test case 2: Authorization header present (fallback)
|
||||
(
|
||||
[(b"authorization", b"Bearer test-auth-token")],
|
||||
"Bearer test-auth-token",
|
||||
"test-auth-token",
|
||||
None,
|
||||
{},
|
||||
),
|
||||
|
|
@ -674,7 +674,9 @@ class TestMCPOAuth2AuthFlow:
|
|||
) = await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
||||
# Should succeed with the LiteLLM key from Authorization header
|
||||
assert auth_result.api_key == "Bearer sk-litellm-valid-key"
|
||||
from litellm.proxy.utils import hash_token
|
||||
|
||||
assert auth_result.api_key == hash_token("sk-litellm-valid-key")
|
||||
mock_auth.assert_called_once()
|
||||
|
||||
async def test_non_auth_http_exception_still_raises(self):
|
||||
|
|
|
|||
|
|
@ -92,6 +92,37 @@ class MockLiteLLMEndUserTable:
|
|||
return self._find_many_results
|
||||
|
||||
|
||||
class MockBatcher:
|
||||
"""Captures per-row update calls and exposes them after commit().
|
||||
|
||||
Mirrors prisma's `db.batch_()` ergonomics enough that the reset job's
|
||||
narrow-write helpers (`_write_key_reset_updates` et al) can run against
|
||||
the mock and the test can assert on what would have been written.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.calls: List[Dict[str, Any]] = []
|
||||
self.committed: bool = False
|
||||
|
||||
class _Table:
|
||||
def __init__(_self, table_name: str, outer: "MockBatcher"):
|
||||
_self._table_name = table_name
|
||||
_self._outer = outer
|
||||
|
||||
def update(_self, where, data):
|
||||
_self._outer.calls.append(
|
||||
{"table": _self._table_name, "where": where, "data": data}
|
||||
)
|
||||
|
||||
self.litellm_verificationtoken = _Table("key", self)
|
||||
self.litellm_usertable = _Table("user", self)
|
||||
self.litellm_teamtable = _Table("team", self)
|
||||
|
||||
async def commit(self):
|
||||
self.committed = True
|
||||
return self.calls
|
||||
|
||||
|
||||
class MockDB:
|
||||
def __init__(self):
|
||||
self.litellm_teammembership = MockLiteLLMTeamMembership()
|
||||
|
|
@ -99,6 +130,19 @@ class MockDB:
|
|||
self.litellm_endusertable = MockLiteLLMEndUserTable()
|
||||
self.litellm_organizationtable = MockLiteLLMOrganizationTable()
|
||||
self.litellm_tagtable = MockLiteLLMTagTable()
|
||||
self.batch_calls: List[Dict[str, Any]] = []
|
||||
|
||||
def batch_(self):
|
||||
batcher = MockBatcher()
|
||||
# Aggregate calls across all batches so tests can assert on cumulative writes.
|
||||
original_commit = batcher.commit
|
||||
|
||||
async def _record_and_commit():
|
||||
self.batch_calls.extend(batcher.calls)
|
||||
return await original_commit()
|
||||
|
||||
batcher.commit = _record_and_commit # type: ignore[assignment]
|
||||
return batcher
|
||||
|
||||
|
||||
class MockPrismaClient:
|
||||
|
|
@ -205,6 +249,7 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
|
|||
"budget_duration": "30d",
|
||||
"budget_reset_at": now,
|
||||
"id": "test-key-1",
|
||||
"token": "tok-key-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -213,11 +258,16 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
|
|||
# Run the test
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
|
||||
|
||||
# Verify results
|
||||
assert len(mock_prisma_client.updated_data["key"]) == 1
|
||||
updated_key = mock_prisma_client.updated_data["key"][0]
|
||||
assert updated_key.spend == 0.0
|
||||
assert updated_key.budget_reset_at > now
|
||||
# The reset writes only {spend, budget_reset_at} per row via batch_().
|
||||
# Full-row writes would re-detonate the Prisma DataError on rows carrying
|
||||
# object_permission_id / budget_limits (see #27730).
|
||||
key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"]
|
||||
assert len(key_writes) == 1
|
||||
write = key_writes[0]
|
||||
assert write["where"] == {"token": "tok-key-1"}
|
||||
assert write["data"]["spend"] == 0
|
||||
assert write["data"]["budget_reset_at"] > now
|
||||
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
|
||||
|
||||
def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
|
||||
|
|
@ -231,6 +281,7 @@ def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
|
|||
"budget_duration": "7d",
|
||||
"budget_reset_at": now,
|
||||
"id": "test-user-1",
|
||||
"user_id": "uid-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -239,11 +290,13 @@ def test_reset_budget_for_user(reset_budget_job, mock_prisma_client):
|
|||
# Run the test
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_users())
|
||||
|
||||
# Verify results
|
||||
assert len(mock_prisma_client.updated_data["user"]) == 1
|
||||
updated_user = mock_prisma_client.updated_data["user"][0]
|
||||
assert updated_user.spend == 0.0
|
||||
assert updated_user.budget_reset_at > now
|
||||
user_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "user"]
|
||||
assert len(user_writes) == 1
|
||||
write = user_writes[0]
|
||||
assert write["where"] == {"user_id": "uid-1"}
|
||||
assert write["data"]["spend"] == 0
|
||||
assert write["data"]["budget_reset_at"] > now
|
||||
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
|
||||
|
||||
def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
|
||||
|
|
@ -257,6 +310,7 @@ def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
|
|||
"budget_duration": "1mo",
|
||||
"budget_reset_at": now,
|
||||
"id": "test-team-1",
|
||||
"team_id": "tid-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -265,11 +319,13 @@ def test_reset_budget_for_team(reset_budget_job, mock_prisma_client):
|
|||
# Run the test
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_teams())
|
||||
|
||||
# Verify results
|
||||
assert len(mock_prisma_client.updated_data["team"]) == 1
|
||||
updated_team = mock_prisma_client.updated_data["team"][0]
|
||||
assert updated_team.spend == 0.0
|
||||
assert updated_team.budget_reset_at > now
|
||||
team_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "team"]
|
||||
assert len(team_writes) == 1
|
||||
write = team_writes[0]
|
||||
assert write["where"] == {"team_id": "tid-1"}
|
||||
assert write["data"]["spend"] == 0
|
||||
assert write["data"]["budget_reset_at"] > now
|
||||
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
|
||||
|
||||
def test_reset_budget_for_enduser(reset_budget_job, mock_prisma_client):
|
||||
|
|
@ -324,6 +380,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
|
|||
"budget_duration": "30d",
|
||||
"budget_reset_at": now,
|
||||
"id": "test-key-1",
|
||||
"token": "tok-all-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -335,6 +392,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
|
|||
"budget_duration": "7d",
|
||||
"budget_reset_at": now,
|
||||
"id": "test-user-1",
|
||||
"user_id": "uid-all-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -346,6 +404,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
|
|||
"budget_duration": "1mo",
|
||||
"budget_reset_at": now,
|
||||
"id": "test-team-1",
|
||||
"team_id": "tid-all-1",
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -379,17 +438,22 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client):
|
|||
# Run the test
|
||||
asyncio.run(reset_budget_job.reset_budget())
|
||||
|
||||
# Verify results
|
||||
assert len(mock_prisma_client.updated_data["key"]) == 1
|
||||
assert len(mock_prisma_client.updated_data["user"]) == 1
|
||||
assert len(mock_prisma_client.updated_data["team"]) == 1
|
||||
# key/user/team rows are written via batch_().<table>.update — verify each
|
||||
# one fired exactly once with the narrow {spend, budget_reset_at} payload.
|
||||
for table_name, where in [
|
||||
("key", {"token": "tok-all-1"}),
|
||||
("user", {"user_id": "uid-all-1"}),
|
||||
("team", {"team_id": "tid-all-1"}),
|
||||
]:
|
||||
writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == table_name]
|
||||
assert len(writes) == 1, f"expected 1 {table_name} write, got {len(writes)}"
|
||||
assert writes[0]["where"] == where
|
||||
assert writes[0]["data"]["spend"] == 0
|
||||
assert set(writes[0]["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
|
||||
# Enduser + budget rows still go through update_data (not narrowed; different path).
|
||||
assert len(mock_prisma_client.updated_data["enduser"]) == 1
|
||||
assert len(mock_prisma_client.updated_data["budget"]) == 1
|
||||
|
||||
# Check that all spends were reset to 0
|
||||
assert mock_prisma_client.updated_data["key"][0].spend == 0.0
|
||||
assert mock_prisma_client.updated_data["user"][0].spend == 0.0
|
||||
assert mock_prisma_client.updated_data["team"][0].spend == 0.0
|
||||
assert mock_prisma_client.updated_data["enduser"][0].spend == 0.0
|
||||
|
||||
|
||||
|
|
@ -1399,6 +1463,105 @@ def test_reset_budget_for_teams_invalidates_redis_counter(
|
|||
)
|
||||
|
||||
|
||||
def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch):
|
||||
"""
|
||||
Regression for #27730 (the bypass-half).
|
||||
|
||||
If the DB write inside the reset job raises (e.g. Prisma DataError on a
|
||||
row carrying object_permission_id or budget_limits), the Redis spend
|
||||
counter MUST NOT be zeroed — that would let get_current_spend admit
|
||||
requests past the cap while the DB row still holds the over-budget
|
||||
spend.
|
||||
|
||||
Pre-fix: _reset_budget_common pre-zeroed the counter before the DB
|
||||
write attempt, opening the bypass window.
|
||||
Post-fix: counter invalidation lives in the caller, AFTER the DB write
|
||||
commits. If the write raises, the post-write invalidation never runs.
|
||||
"""
|
||||
counter_cache = _make_counter_invalidation_job(monkeypatch)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
prisma_client = MagicMock()
|
||||
|
||||
matching_key = type(
|
||||
"Key",
|
||||
(),
|
||||
{
|
||||
"spend": 100.0,
|
||||
"budget_duration": "30d",
|
||||
"budget_reset_at": now - timedelta(seconds=1),
|
||||
"token": "sk-failing",
|
||||
},
|
||||
)
|
||||
|
||||
# get_data returns one key needing reset; the batched DB write then explodes.
|
||||
async def fake_get_data(table_name, query_type, **kwargs):
|
||||
if table_name == "key":
|
||||
return [matching_key]
|
||||
return []
|
||||
|
||||
prisma_client.get_data = fake_get_data
|
||||
|
||||
batcher = MagicMock()
|
||||
batcher.litellm_verificationtoken.update = MagicMock()
|
||||
|
||||
async def failing_commit():
|
||||
raise RuntimeError("simulated Prisma DataError on update")
|
||||
|
||||
batcher.commit = failing_commit
|
||||
prisma_client.db.batch_ = MagicMock(return_value=batcher)
|
||||
|
||||
job = ResetBudgetJob(
|
||||
proxy_logging_obj=MockProxyLogging(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
asyncio.run(job.reset_budget_for_litellm_keys())
|
||||
|
||||
# CRITICAL: counter invalidation must NOT have been called at all —
|
||||
# the DB write raised before the post-write invalidation loop. Using
|
||||
# assert_not_called() instead of iterating call_args_list, because the
|
||||
# latter is vacuously true when the list is empty (would pass even if
|
||||
# the bypass were re-introduced via a different code path).
|
||||
counter_cache.in_memory_cache.set_cache.assert_not_called()
|
||||
|
||||
|
||||
def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job, mock_prisma_client):
|
||||
"""
|
||||
Regression for #27730 (the trigger-half).
|
||||
|
||||
The reset job must write only {spend, budget_reset_at} per row — never
|
||||
the full key object. Sending the full object via the old update_data
|
||||
batcher path made Prisma reject any row carrying object_permission_id
|
||||
or budget_limits (both became non-NULL on UI-created keys after v1.84.0).
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
key_with_problematic_fields = type(
|
||||
"LiteLLM_VerificationToken",
|
||||
(),
|
||||
{
|
||||
"spend": 50.0,
|
||||
"budget_duration": "30d",
|
||||
"budget_reset_at": now,
|
||||
"token": "sk-problematic",
|
||||
"object_permission_id": "perm-abc", # would be rejected on update
|
||||
"budget_limits": [{"max_budget": 5}], # would be rejected on update
|
||||
"metadata": {"some": "thing"},
|
||||
},
|
||||
)
|
||||
mock_prisma_client.data["key"] = [key_with_problematic_fields]
|
||||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
|
||||
|
||||
key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"]
|
||||
assert len(key_writes) == 1
|
||||
payload_keys = set(key_writes[0]["data"].keys())
|
||||
assert payload_keys == {"spend", "budget_reset_at"}, (
|
||||
f"reset payload must not include any field besides spend / budget_reset_at, "
|
||||
f"got: {payload_keys}. Any extra field (object_permission_id, budget_limits, etc.) "
|
||||
f"trips Prisma DataError and detonates the whole batch."
|
||||
)
|
||||
|
||||
|
||||
def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monkeypatch):
|
||||
"""Resetting keys via budget tier must clear each linked key's counter."""
|
||||
counter_cache = _make_counter_invalidation_job(monkeypatch)
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ import asyncio
|
|||
import os
|
||||
import sys
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -38,6 +38,24 @@ from litellm.types.guardrails import GuardrailEventHooks
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _attach_mock_success_dispatch(mock_logging_obj, async_success_fn):
|
||||
"""Match production entrypoint: ``_run_deferred_stream_guardrails`` uses dispatch."""
|
||||
|
||||
async def dispatch_success_handlers(
|
||||
result=None, start_time=None, end_time=None, cache_hit=None, **kwargs
|
||||
):
|
||||
await async_success_fn(
|
||||
result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
mock_logging_obj.dispatch_success_handlers = dispatch_success_handlers
|
||||
mock_logging_obj.async_success_handler = async_success_fn
|
||||
|
||||
|
||||
class PostCallGuardrail(CustomGuardrail):
|
||||
"""A post-call guardrail."""
|
||||
|
||||
|
|
@ -454,7 +472,7 @@ class TestDeferredStreamingClosure:
|
|||
async def track_async_success(*args, **kwargs):
|
||||
pass
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
tracking_guardrail = TrackingGuardrail()
|
||||
tracking_logger = TrackingLogger()
|
||||
|
|
@ -511,7 +529,7 @@ class TestDeferredStreamingClosure:
|
|||
nonlocal logged_response
|
||||
logged_response = args[0] if args else None
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
class ModifyingGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
|
|
@ -573,7 +591,7 @@ class TestDeferredStreamingClosure:
|
|||
nonlocal logging_called
|
||||
logging_called = True
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
guardrail = BlockingGuardrail()
|
||||
|
||||
|
|
@ -621,7 +639,7 @@ class TestDeferredStreamingClosure:
|
|||
async def track_async_success(*args, **kwargs):
|
||||
pass
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
guardrail = TransientErrorGuardrail()
|
||||
|
||||
|
|
@ -656,7 +674,7 @@ class TestDeferredStreamingClosure:
|
|||
nonlocal logged_response
|
||||
logged_response = args[0] if args else None
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
class TestGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
|
|
@ -739,7 +757,7 @@ class TestDeferredStreamingClosure:
|
|||
async def track_async_success(*args, **kwargs):
|
||||
pass
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
guardrail = ApplyGuardrailType()
|
||||
|
||||
|
|
@ -792,7 +810,7 @@ class TestDeferredStreamingClosure:
|
|||
async def track_async_success(*args, **kwargs):
|
||||
pass
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
guardrail = IteratorHookGuardrail()
|
||||
|
||||
|
|
@ -847,7 +865,7 @@ class TestDeferredStreamingClosure:
|
|||
async def track_async_success(*args, **kwargs):
|
||||
pass
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
guardrail = InspectingGuardrail()
|
||||
|
||||
|
|
@ -914,7 +932,7 @@ class TestDeferredStreamingClosure:
|
|||
async def track_async_success(*args, **kwargs):
|
||||
pass
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
guardrail_a = TaggedGuardrail("guardrail-a")
|
||||
guardrail_b = TaggedGuardrail("guardrail-b")
|
||||
|
|
@ -962,7 +980,7 @@ class TestDeferredStreamingClosure:
|
|||
nonlocal logging_called
|
||||
logging_called = True
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
def exploding_merge(data, llm_router):
|
||||
raise RuntimeError("Simulated init failure")
|
||||
|
|
@ -986,6 +1004,67 @@ class TestDeferredStreamingClosure:
|
|||
logging_called is True
|
||||
), "Logging must fire even when guardrail initialization raises"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deferred_logging_forces_async_for_sync_classified_call_type(self):
|
||||
"""
|
||||
Regression: proxy deferred streaming logging must reach the async success
|
||||
handler (which runs the async-only DB/spend logger) even when the call
|
||||
type is classified as a sync SDK request by _is_sync_litellm_request.
|
||||
|
||||
Without prefer_async_handlers=True, an async proxy stream whose
|
||||
litellm_params lacks a recognized async marker would enter the sync
|
||||
branch of dispatch_success_handlers and silently skip spend tracking.
|
||||
|
||||
Uses the real dispatch_success_handlers via the production
|
||||
_run_deferred_stream_guardrails entrypoint.
|
||||
"""
|
||||
import time
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as LiteLLMLoggingObj,
|
||||
)
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion", # not pass_through_endpoint
|
||||
start_time=time.time(),
|
||||
litellm_call_id="test-id",
|
||||
function_id="fn",
|
||||
)
|
||||
# litellm_params with no recognized async marker -> classified sync.
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
assert LiteLLMLoggingObj._is_sync_litellm_request({}) is True
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
logging_obj, "async_success_handler", new_callable=AsyncMock
|
||||
) as mock_async,
|
||||
patch.object(
|
||||
logging_obj, "success_handler", new_callable=MagicMock
|
||||
) as mock_sync,
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_callbacks_for_async_calls",
|
||||
return_value=False,
|
||||
),
|
||||
patch("litellm.callbacks", [PostCallGuardrail()]),
|
||||
):
|
||||
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
|
||||
captured_data={"model": "gpt-4o-mini", "metadata": {}},
|
||||
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
|
||||
captured_logging_obj=logging_obj,
|
||||
assembled_response=MagicMock(),
|
||||
cache_hit=False,
|
||||
)
|
||||
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
mock_async.assert_awaited_once()
|
||||
mock_sync.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 7. _fire_deferred_stream_logging
|
||||
|
|
@ -1054,7 +1133,7 @@ class TestFireDeferredStreamLogging:
|
|||
nonlocal logged_response
|
||||
logged_response = args[0] if args else None
|
||||
|
||||
mock_logging_obj.async_success_handler = track_async_success
|
||||
_attach_mock_success_dispatch(mock_logging_obj, track_async_success)
|
||||
|
||||
class InfoWritingGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
|
|
|
|||
|
|
@ -2775,3 +2775,121 @@ async def test_project_model_rate_limits_not_triggered_for_other_model_v3():
|
|||
assert (
|
||||
"model_per_project" not in descriptor_keys
|
||||
), f"model_per_project should not be added for unrelated model, got: {descriptor_keys}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_does_not_leak_internal_stash_to_request_body():
|
||||
"""Regression for #27001: stash keys must stay in metadata, never on
|
||||
the top level of ``data`` (which gets forwarded as the provider body)."""
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_LITELLM_STASH_KEYS,
|
||||
RATE_LIMIT_DESCRIPTORS_KEY,
|
||||
TPM_RESERVED_TOKENS_KEY,
|
||||
)
|
||||
|
||||
_api_key = hash_token("sk-leak-regression")
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key=_api_key,
|
||||
tpm_limit=1000,
|
||||
rpm_limit=5,
|
||||
)
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache),
|
||||
)
|
||||
|
||||
async def mock_should_rate_limit(descriptors, **kwargs):
|
||||
return {"overall_code": "OK", "statuses": []}
|
||||
|
||||
async def mock_reserve_tpm_tokens(descriptors, estimated_tokens, **kwargs):
|
||||
return {
|
||||
"overall_code": "OK",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OK",
|
||||
"current_limit": 1000,
|
||||
"limit_remaining": 1000 - estimated_tokens,
|
||||
"descriptor_key": d["key"],
|
||||
"descriptor_value": d["value"],
|
||||
"rate_limit_type": "tokens",
|
||||
}
|
||||
for d in descriptors
|
||||
],
|
||||
}
|
||||
|
||||
parallel_request_handler.should_rate_limit = mock_should_rate_limit
|
||||
parallel_request_handler.reserve_tpm_tokens = mock_reserve_tpm_tokens
|
||||
|
||||
data: Dict[str, Any] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 10,
|
||||
}
|
||||
|
||||
await parallel_request_handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
leaked = [k for k in _LITELLM_STASH_KEYS if k in data]
|
||||
assert not leaked, f"stash keys leaked to top level: {leaked}"
|
||||
|
||||
metadata = data.get("metadata") or {}
|
||||
assert metadata.get(TPM_RESERVED_TOKENS_KEY)
|
||||
assert isinstance(metadata.get(RATE_LIMIT_DESCRIPTORS_KEY), list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_rejects_caller_supplied_stash_values():
|
||||
"""Caller cannot pre-populate stash keys in body metadata to drive a
|
||||
later TPM refund against an arbitrary scope."""
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_LITELLM_STASH_KEYS,
|
||||
RATE_LIMIT_DESCRIPTORS_KEY,
|
||||
TPM_RESERVED_TOKENS_KEY,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-no-limits"))
|
||||
local_cache = DualCache()
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache),
|
||||
)
|
||||
|
||||
victim_descriptors = [
|
||||
{
|
||||
"key": "api_key",
|
||||
"value": "victim-key-hash",
|
||||
"rate_limit": {"tokens_per_unit": 10000, "window_size": 60},
|
||||
}
|
||||
]
|
||||
data: Dict[str, Any] = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
TPM_RESERVED_TOKENS_KEY: 9999,
|
||||
RATE_LIMIT_DESCRIPTORS_KEY: victim_descriptors,
|
||||
"metadata": {
|
||||
TPM_RESERVED_TOKENS_KEY: 9999,
|
||||
RATE_LIMIT_DESCRIPTORS_KEY: victim_descriptors,
|
||||
},
|
||||
"litellm_metadata": {
|
||||
TPM_RESERVED_TOKENS_KEY: 9999,
|
||||
RATE_LIMIT_DESCRIPTORS_KEY: victim_descriptors,
|
||||
},
|
||||
}
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
for channel in (
|
||||
data,
|
||||
data.get("metadata") or {},
|
||||
data.get("litellm_metadata") or {},
|
||||
):
|
||||
leaked = [k for k in _LITELLM_STASH_KEYS if k in channel]
|
||||
assert not leaked, f"caller-supplied stash survived in {channel!r}: {leaked}"
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ import pytest
|
|||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
RATE_LIMIT_DESCRIPTORS_KEY,
|
||||
TPM_RESERVATION_RELEASED_KEY,
|
||||
TPM_RESERVED_MODEL_KEY,
|
||||
TPM_RESERVED_SCOPES_KEY,
|
||||
|
|
@ -606,9 +607,9 @@ async def test_contentless_request_reserves_minimum(rate_limiter):
|
|||
data=data,
|
||||
call_type="",
|
||||
)
|
||||
assert (
|
||||
data.get(TPM_RESERVED_TOKENS_KEY) == 1
|
||||
), "Contentless request should reserve the floor of 1 token"
|
||||
assert (data.get("metadata") or {}).get(
|
||||
TPM_RESERVED_TOKENS_KEY
|
||||
) == 1, "Contentless request should reserve the floor of 1 token"
|
||||
|
||||
counter_after_two = int(
|
||||
await cache.async_get_cache(key=counter_key, local_only=True) or 0
|
||||
|
|
@ -701,7 +702,7 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter):
|
|||
data=data,
|
||||
call_type="",
|
||||
)
|
||||
reserved = data[TPM_RESERVED_TOKENS_KEY]
|
||||
reserved = (data.get("metadata") or {})[TPM_RESERVED_TOKENS_KEY]
|
||||
assert reserved > 0
|
||||
|
||||
counter_key = handler.create_rate_limit_keys(
|
||||
|
|
@ -726,9 +727,9 @@ async def test_reservation_released_on_proxy_rejection(rate_limiter):
|
|||
f"Reservation leaked: counter={counter_after_release} after "
|
||||
f"proxy-level rejection refund (expected 0)."
|
||||
)
|
||||
assert data.get(TPM_RESERVATION_RELEASED_KEY) is True, (
|
||||
"Released marker must be stamped to prevent async_log_failure_event "
|
||||
"from double-refunding."
|
||||
assert (data.get("metadata") or {}).get(TPM_RESERVATION_RELEASED_KEY) is True, (
|
||||
"Released marker must be stamped to prevent "
|
||||
"async_log_failure_event from double-refunding."
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -760,12 +761,7 @@ async def test_reservation_release_idempotent(rate_limiter):
|
|||
shared_metadata = {
|
||||
"user_api_key_hash": api_key,
|
||||
TPM_RESERVED_TOKENS_KEY: 100,
|
||||
}
|
||||
|
||||
request_data = {
|
||||
"metadata": shared_metadata,
|
||||
TPM_RESERVED_TOKENS_KEY: 100,
|
||||
"_litellm_rate_limit_descriptors": [
|
||||
RATE_LIMIT_DESCRIPTORS_KEY: [
|
||||
{
|
||||
"key": "api_key",
|
||||
"value": api_key,
|
||||
|
|
@ -774,6 +770,10 @@ async def test_reservation_release_idempotent(rate_limiter):
|
|||
],
|
||||
}
|
||||
|
||||
request_data = {
|
||||
"metadata": shared_metadata,
|
||||
}
|
||||
|
||||
await handler.async_post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=Exception("rejected"),
|
||||
|
|
|
|||
|
|
@ -69,3 +69,21 @@ def test_internal_jobs_user_has_proxy_admin_role():
|
|||
assert system_user.user_id == "system"
|
||||
assert system_user.team_id == "system"
|
||||
assert system_user.team_alias == "system"
|
||||
|
||||
|
||||
def test_user_api_key_auth_hashes_authorization_header_form_of_key():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
raw_key = "sk-AbCdEfGhIjKlMnOpQrStUvWxYz0123456789"
|
||||
baseline = UserAPIKeyAuth(api_key=raw_key)
|
||||
|
||||
for header_form in (
|
||||
f"Bearer {raw_key}",
|
||||
f"bearer {raw_key}",
|
||||
f"BEARER {raw_key}",
|
||||
f"BeArEr {raw_key}",
|
||||
):
|
||||
from_header = UserAPIKeyAuth(api_key=header_form)
|
||||
assert from_header.api_key == baseline.api_key
|
||||
assert from_header.token == baseline.token
|
||||
assert not from_header.api_key.lower().startswith("bearer")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue