mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(otel): nest cache spans under their operation and name service spans by purpose (#44150)
Response cache reads and writes open cache.get llm_response and cache.set llm_response phase spans with their Redis spans nested underneath, on the Python path and on the native Rust path, and deployment selection runs inside a route {model_group} phase so the cooldown, usage and model-id reads the router issues nest under it before chat {model}. The autorouter classifier call nests under that route phase as well and carries its typed internal origin on litellm.request.purpose, so it is told apart from the provider attempt. Service spans are named {service}.{verb} {target} from a low-cardinality key family the producer declares (llm_response, auth_objects, spend_counters, router_cooldowns, claude_code_session_router_binding, rate_limits, pod_lock, budget_reset, ...) instead of the raw method or a per-request pipeline length; a pipeline flush is targeted by the one family its ops share or by mixed with the sorted families on litellm.redis.families, a batch op keeps the family it was declared under whichever pipeline or standalone read settles it, and the ambient family labels Redis spans only, never the DB write-back a task spawned inside that context performs later. The raw method stays on litellm.service.call_type and on the Prometheus and Datadog labels. Caller attribution is carried across asyncio task boundaries on a ContextVar so forwarder-only chains no longer surface, the raw cache key is dropped from Redis span metadata, pipeline op counts land as an integer attribute, every call_type the Redis cache layer emits maps to a verb, and a scan over litellm/ and enterprise/ fails when a Redis producer, batch reservation included, declares no key family.
A V2 logger built for a key or team logging entry while the operator's V2 logger is already registered keeps only the exporters its own preset contributed, whether or not the operator holds credentials for that backend, so every chat span no longer reaches the operator's collector twice. A span the success callback has to open itself, with no pre-call carrier, starts at the provider handoff (api_call_start_time) instead of the logging object's creation.
Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
53d2ab6b05
commit
564d236985
114 changed files with 3070 additions and 987 deletions
|
|
@ -7,15 +7,19 @@
|
|||
## This accepts a list of user id's for whom calls will be rejected
|
||||
|
||||
|
||||
from typing import Optional, Literal
|
||||
import litellm
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth, LiteLLM_EndUserTable
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from typing import Literal, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import LiteLLM_EndUserTable, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class _ENTERPRISE_BlockedUserList(CustomLogger):
|
||||
enforces_request_content: bool = True
|
||||
|
|
@ -54,6 +58,7 @@ class _ENTERPRISE_BlockedUserList(CustomLogger):
|
|||
if litellm.set_verbose is True:
|
||||
print(print_statement) # noqa
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Base class for sending emails to user after creating keys or invite links
|
|||
import html
|
||||
import json
|
||||
import os
|
||||
from typing import List, Literal, Optional
|
||||
from typing import Final, List, Literal, Optional
|
||||
|
||||
from litellm_enterprise.types.enterprise_callbacks.send_emails import (
|
||||
EmailEvent,
|
||||
|
|
@ -15,6 +15,7 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import (
|
|||
SendKeyRotatedEmailEvent,
|
||||
)
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -48,6 +49,8 @@ from litellm.proxy._types import (
|
|||
from litellm.secret_managers.main import get_secret_bool
|
||||
from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL
|
||||
|
||||
_BUDGET_ALERT_CLAIMS_TARGET: Final = "budget_alert_claims"
|
||||
|
||||
|
||||
def _max_budget_alert_id(user_info: CallInfo) -> str:
|
||||
if user_info.event_group == Litellm_EntityType.TEAM_MEMBER:
|
||||
|
|
@ -437,6 +440,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
html_body=email_html_content,
|
||||
)
|
||||
|
||||
@with_service_target(_BUDGET_ALERT_CLAIMS_TARGET)
|
||||
async def budget_alerts(
|
||||
self,
|
||||
type: Literal[
|
||||
|
|
@ -606,6 +610,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
await self._release_budget_alert_claim(_cache, _cache_key)
|
||||
return
|
||||
|
||||
@with_service_target(_BUDGET_ALERT_CLAIMS_TARGET)
|
||||
async def _handle_multi_threshold_max_budget_alert(
|
||||
self,
|
||||
user_info: CallInfo,
|
||||
|
|
@ -691,6 +696,7 @@ class BaseEmailLogger(CustomLogger):
|
|||
)
|
||||
await self._release_budget_alert_claim(_cache, _cache_key)
|
||||
|
||||
@with_service_target(_BUDGET_ALERT_CLAIMS_TARGET)
|
||||
async def _release_budget_alert_claim(self, cache: DualCache, cache_key: str) -> None:
|
||||
try:
|
||||
await cache.async_delete_cache(key=cache_key)
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from pydantic import ValidationError
|
|||
|
||||
import litellm
|
||||
from litellm import Router, verbose_logger
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import MAX_FILE_LIST_LIMIT
|
||||
|
|
@ -229,6 +230,9 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s
|
|||
)
|
||||
|
||||
|
||||
_MANAGED_FILES_TARGET: Final = "managed_files"
|
||||
|
||||
|
||||
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
||||
# Class variables or attributes
|
||||
def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
|
||||
|
|
@ -242,6 +246,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
|
||||
return PrometheusLogger.get_instance()
|
||||
|
||||
@with_service_target(_MANAGED_FILES_TARGET)
|
||||
async def store_unified_file_id(
|
||||
self,
|
||||
file_id: str,
|
||||
|
|
@ -325,6 +330,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
verbose_logger.warning(f"could not resolve org for managed object attribution: {e}")
|
||||
return None
|
||||
|
||||
@with_service_target(_MANAGED_FILES_TARGET)
|
||||
async def store_unified_object_id(
|
||||
self,
|
||||
unified_object_id: str,
|
||||
|
|
@ -412,6 +418,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
},
|
||||
)
|
||||
|
||||
@with_service_target(_MANAGED_FILES_TARGET)
|
||||
async def get_unified_file_id(
|
||||
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
|
||||
) -> Optional[LiteLLM_ManagedFileTable]:
|
||||
|
|
@ -434,6 +441,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
return LiteLLM_ManagedFileTable.model_validate(db_object.model_dump())
|
||||
return None
|
||||
|
||||
@with_service_target(_MANAGED_FILES_TARGET)
|
||||
async def delete_unified_file_id(
|
||||
self, file_id: str, litellm_parent_otel_span: Optional[Span] = None
|
||||
) -> OpenAIFileObject:
|
||||
|
|
|
|||
|
|
@ -6,11 +6,17 @@ be settable from user input. Context variables are scoped to the current
|
|||
asyncio task and cannot be injected via HTTP request bodies.
|
||||
"""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
import inspect
|
||||
from collections.abc import Awaitable, Callable, Generator
|
||||
from contextlib import contextmanager, suppress
|
||||
from contextvars import ContextVar, Token
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
from functools import wraps
|
||||
from typing import Final, ParamSpec, TypeVar, cast
|
||||
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
_T = TypeVar("_T")
|
||||
|
||||
# When True, suppresses async logging and billing for internal sub-calls
|
||||
# (e.g., emulated file-search steps that make nested LLM calls).
|
||||
|
|
@ -23,6 +29,13 @@ _billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", d
|
|||
|
||||
_post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False)
|
||||
|
||||
_service_target: Final[ContextVar[str | None]] = ContextVar("service_target", default=None)
|
||||
# Event-metadata key under which a Redis pipeline reports the sorted, comma-joined families its ops were
|
||||
# declared under when they span more than one.
|
||||
REDIS_FAMILIES_METADATA_KEY: Final = "families"
|
||||
|
||||
_service_caller: Final[ContextVar[str | None]] = ContextVar("service_caller", default=None)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def post_response_phase() -> Generator[None]:
|
||||
|
|
@ -38,6 +51,68 @@ def in_post_response_phase() -> bool:
|
|||
return _post_response.get()
|
||||
|
||||
|
||||
def _restore(var: ContextVar[_T], token: Token[_T]) -> None:
|
||||
"""Reset ``var``; a coroutine the GC closes from another context has no value left to restore."""
|
||||
with suppress(ValueError):
|
||||
var.reset(token)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def service_target(target: str | None) -> Generator[None]:
|
||||
"""Name what the datastore calls inside this block are for; ``None`` clears an inherited target."""
|
||||
token: Final = _service_target.set(target)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_restore(_service_target, token)
|
||||
|
||||
|
||||
def current_service_target() -> str | None:
|
||||
return _service_target.get()
|
||||
|
||||
|
||||
def with_service_target(target: str) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]:
|
||||
"""Run every call of the decorated function, coroutine functions included, under ``service_target(target)``."""
|
||||
|
||||
def decorate(fn: Callable[_P, _R]) -> Callable[_P, _R]:
|
||||
if inspect.iscoroutinefunction(fn):
|
||||
awaitable_fn: Final[Callable[_P, Awaitable[object]]] = cast( # cast-ok: checked by iscoroutinefunction
|
||||
"Callable[_P, Awaitable[object]]", fn
|
||||
)
|
||||
|
||||
@wraps(fn)
|
||||
async def run_async(*args: _P.args, **kwargs: _P.kwargs) -> object:
|
||||
with service_target(target):
|
||||
return await awaitable_fn(*args, **kwargs)
|
||||
|
||||
return cast("Callable[_P, _R]", run_async) # cast-ok: same coroutine-returning signature as ``fn``
|
||||
|
||||
@wraps(fn)
|
||||
def run(*args: _P.args, **kwargs: _P.kwargs) -> _R:
|
||||
with service_target(target):
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
return run
|
||||
|
||||
return decorate
|
||||
|
||||
|
||||
@contextmanager
|
||||
def service_caller(caller: str | None) -> Generator[None]:
|
||||
"""Name the litellm code a datastore call was issued for when its own frames cannot: an operation
|
||||
declared in one task and run in another (a batch op retried on the flush) carries the chain captured
|
||||
where it was declared."""
|
||||
token: Final = _service_caller.set(caller)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_restore(_service_caller, token)
|
||||
|
||||
|
||||
def current_service_caller() -> str | None:
|
||||
return _service_caller.get()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def pinned_billing_time(moment: datetime) -> Generator[None]:
|
||||
"""Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from datetime import datetime, timedelta
|
|||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
from .integrations.custom_logger import CustomLogger
|
||||
|
|
@ -234,6 +235,7 @@ class ServiceLogging(CustomLogger):
|
|||
duration=duration,
|
||||
call_type=call_type,
|
||||
caller=caller,
|
||||
target=current_service_target() if service == ServiceTypes.REDIS else None,
|
||||
event_metadata=event_metadata,
|
||||
)
|
||||
|
||||
|
|
@ -340,6 +342,7 @@ class ServiceLogging(CustomLogger):
|
|||
duration=duration,
|
||||
call_type=call_type,
|
||||
caller=caller,
|
||||
target=current_service_target() if service == ServiceTypes.REDIS else None,
|
||||
event_metadata=event_metadata,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ from pydantic import JsonValue, TypeAdapter, ValidationError
|
|||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
ROUTER_SESSION_PINS_TARGET: Final = "router_session_pins"
|
||||
|
||||
_PIN_JSON_ADAPTER: Final = TypeAdapter[JsonValue](JsonValue)
|
||||
|
||||
_CLAIM_PIN_SCRIPT: Final = """
|
||||
|
|
|
|||
|
|
@ -13,16 +13,19 @@ import json
|
|||
import logging
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_service_target, service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import CACHED_STREAMING_CHUNK_DELAY
|
||||
from litellm.integrations.otel.runtime import phase_span
|
||||
from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
|
||||
from litellm.types.caching import *
|
||||
from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg
|
||||
|
|
@ -59,6 +62,21 @@ def _native_response(result: object) -> object:
|
|||
return result
|
||||
|
||||
|
||||
RESPONSE_CACHE_TARGET: Final = "llm_response"
|
||||
|
||||
|
||||
@contextmanager
|
||||
def response_cache_phase(operation: Literal["get", "set"]) -> Generator[None]:
|
||||
"""The ``cache.get llm_response`` / ``cache.set llm_response`` span a response-cache read or write runs
|
||||
inside, so its datastore spans nest under it and read by purpose. Entered by the facade methods so every
|
||||
caller gets it (the native bridge calls them straight); a call already inside the phase keeps it."""
|
||||
if current_service_target() == RESPONSE_CACHE_TARGET:
|
||||
yield
|
||||
return
|
||||
with phase_span(f"cache.{operation} {RESPONSE_CACHE_TARGET}"), service_target(RESPONSE_CACHE_TARGET):
|
||||
yield
|
||||
|
||||
|
||||
def print_verbose(print_statement):
|
||||
try:
|
||||
verbose_logger.debug(print_statement)
|
||||
|
|
@ -615,32 +633,33 @@ class Cache:
|
|||
try: # never block execution
|
||||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
if "cache_key" in kwargs:
|
||||
cache_key = kwargs["cache_key"]
|
||||
else:
|
||||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None and self._native_cache is not None:
|
||||
request = self._native_cache.request(self, MappingProxyType({**kwargs, "cache_key": cache_key}))
|
||||
if request is None:
|
||||
return None
|
||||
if not self._is_semantic_cache():
|
||||
return self._native_cache.lookup(request)
|
||||
response, similarity = self._native_cache.lookup_semantic(request)
|
||||
self._stamp_semantic_similarity(kwargs, similarity)
|
||||
return response
|
||||
if cache_key is not None:
|
||||
cache_control_args: Final[DynamicCacheControl] = kwargs.get("cache", {})
|
||||
max_age = cache_control_args.get("s-maxage") or cache_control_args.get("s-max-age") or float("inf")
|
||||
cache_lookup_kwargs: Final = self._get_safe_cache_lookup_kwargs(kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = dynamic_cache_object.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
with response_cache_phase("get"):
|
||||
if "cache_key" in kwargs:
|
||||
cache_key = kwargs["cache_key"]
|
||||
else:
|
||||
cached_result = self.cache.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
self._update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs=kwargs,
|
||||
cache_lookup_kwargs=cache_lookup_kwargs,
|
||||
)
|
||||
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None and self._native_cache is not None:
|
||||
request = self._native_cache.request(self, MappingProxyType({**kwargs, "cache_key": cache_key}))
|
||||
if request is None:
|
||||
return None
|
||||
if not self._is_semantic_cache():
|
||||
return self._native_cache.lookup(request)
|
||||
response, similarity = self._native_cache.lookup_semantic(request)
|
||||
self._stamp_semantic_similarity(kwargs, similarity)
|
||||
return response
|
||||
if cache_key is not None:
|
||||
cache_control_args: Final[DynamicCacheControl] = kwargs.get("cache", {})
|
||||
max_age = cache_control_args.get("s-maxage") or cache_control_args.get("s-max-age") or float("inf")
|
||||
cache_lookup_kwargs: Final = self._get_safe_cache_lookup_kwargs(kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = dynamic_cache_object.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
else:
|
||||
cached_result = self.cache.get_cache(cache_key, **cache_lookup_kwargs)
|
||||
self._update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs=kwargs,
|
||||
cache_lookup_kwargs=cache_lookup_kwargs,
|
||||
)
|
||||
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
except Exception:
|
||||
print_verbose(f"An exception occurred: {traceback.format_exc()}")
|
||||
return None
|
||||
|
|
@ -656,27 +675,30 @@ class Cache:
|
|||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
|
||||
if "cache_key" in kwargs:
|
||||
cache_key = kwargs["cache_key"]
|
||||
else:
|
||||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None and self._native_cache is not None:
|
||||
request = self._native_cache.request(self, MappingProxyType({**kwargs, "cache_key": cache_key}))
|
||||
if request is None:
|
||||
return None
|
||||
if not self._is_semantic_cache():
|
||||
return await self._native_cache.async_lookup(request)
|
||||
response, similarity = await self._native_cache.async_lookup_semantic(request)
|
||||
self._stamp_semantic_similarity(kwargs, similarity)
|
||||
return response
|
||||
if cache_key is not None:
|
||||
cache_control_args: Final = kwargs.get("cache", {})
|
||||
max_age: Final = cache_control_args.get("s-max-age", cache_control_args.get("s-maxage", float("inf")))
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs)
|
||||
with response_cache_phase("get"):
|
||||
if "cache_key" in kwargs:
|
||||
cache_key = kwargs["cache_key"]
|
||||
else:
|
||||
cached_result = await self.cache.async_get_cache(cache_key, **kwargs)
|
||||
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
cache_key = self.get_cache_key(**kwargs)
|
||||
if cache_key is not None and self._native_cache is not None:
|
||||
request = self._native_cache.request(self, MappingProxyType({**kwargs, "cache_key": cache_key}))
|
||||
if request is None:
|
||||
return None
|
||||
if not self._is_semantic_cache():
|
||||
return await self._native_cache.async_lookup(request)
|
||||
response, similarity = await self._native_cache.async_lookup_semantic(request)
|
||||
self._stamp_semantic_similarity(kwargs, similarity)
|
||||
return response
|
||||
if cache_key is not None:
|
||||
cache_control_args: Final = kwargs.get("cache", {})
|
||||
max_age: Final = cache_control_args.get(
|
||||
"s-max-age", cache_control_args.get("s-maxage", float("inf"))
|
||||
)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = await dynamic_cache_object.async_get_cache(cache_key, **kwargs)
|
||||
else:
|
||||
cached_result = await self.cache.async_get_cache(cache_key, **kwargs)
|
||||
return self._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
except Exception:
|
||||
print_verbose(f"An exception occurred: {traceback.format_exc()}")
|
||||
return None
|
||||
|
|
@ -725,13 +747,14 @@ class Cache:
|
|||
try:
|
||||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
if self._native_cache is not None:
|
||||
request = self._native_request(kwargs)
|
||||
if request is not None:
|
||||
self._native_cache.store(request, _native_response(result))
|
||||
return
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
self.cache.set_cache(cache_key, cached_data, **kwargs)
|
||||
with response_cache_phase("set"):
|
||||
if self._native_cache is not None:
|
||||
request = self._native_request(kwargs)
|
||||
if request is not None:
|
||||
self._native_cache.store(request, _native_response(result))
|
||||
return
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
self.cache.set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
self._log_add_cache_failure(e)
|
||||
|
||||
|
|
@ -749,20 +772,21 @@ class Cache:
|
|||
try:
|
||||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
if self._native_cache is not None:
|
||||
request = self._native_request(kwargs)
|
||||
if request is not None:
|
||||
await self._native_cache.async_store(request, _native_response(result))
|
||||
return
|
||||
if self.type == "redis" and self.redis_flush_size is not None:
|
||||
# high traffic - fill in results in memory and then flush
|
||||
await self.batch_cache_write(result, **kwargs)
|
||||
else:
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
with response_cache_phase("set"):
|
||||
if self._native_cache is not None:
|
||||
request = self._native_request(kwargs)
|
||||
if request is not None:
|
||||
await self._native_cache.async_store(request, _native_response(result))
|
||||
return
|
||||
if self.type == "redis" and self.redis_flush_size is not None:
|
||||
# high traffic - fill in results in memory and then flush
|
||||
await self.batch_cache_write(result, **kwargs)
|
||||
else:
|
||||
await self.cache.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
else:
|
||||
await self.cache.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
self._log_add_cache_failure(e)
|
||||
|
||||
|
|
@ -909,47 +933,50 @@ class Cache:
|
|||
if self.should_use_cache(**kwargs) is not True:
|
||||
return
|
||||
|
||||
input_count: Final = len(kwargs["input"]) if isinstance(kwargs["input"], list) else 1
|
||||
if len(result.data) != input_count:
|
||||
verbose_logger.debug(
|
||||
"LiteLLM Cache: skipping embedding cache write, %d inputs but %d embeddings in the response",
|
||||
input_count,
|
||||
len(result.data),
|
||||
)
|
||||
return
|
||||
with response_cache_phase("set"):
|
||||
input_count: Final = len(kwargs["input"]) if isinstance(kwargs["input"], list) else 1
|
||||
if len(result.data) != input_count:
|
||||
verbose_logger.debug(
|
||||
"LiteLLM Cache: skipping embedding cache write, %d inputs but %d embeddings in the response",
|
||||
input_count,
|
||||
len(result.data),
|
||||
)
|
||||
return
|
||||
|
||||
# set default ttl if not set
|
||||
if self.ttl is not None:
|
||||
kwargs["ttl"] = self.ttl
|
||||
# set default ttl if not set
|
||||
if self.ttl is not None:
|
||||
kwargs["ttl"] = self.ttl
|
||||
|
||||
cache_list: Final = []
|
||||
if isinstance(kwargs["input"], list):
|
||||
for idx, i in enumerate(kwargs["input"]):
|
||||
(
|
||||
cache_key,
|
||||
cached_data,
|
||||
kwargs,
|
||||
) = self.add_embedding_response_to_cache(result, i, kwargs, idx)
|
||||
cache_list: Final = []
|
||||
if isinstance(kwargs["input"], list):
|
||||
for idx, i in enumerate(kwargs["input"]):
|
||||
(
|
||||
cache_key,
|
||||
cached_data,
|
||||
kwargs,
|
||||
) = self.add_embedding_response_to_cache(result, i, kwargs, idx)
|
||||
cache_list.append((cache_key, cached_data))
|
||||
elif isinstance(kwargs["input"], str):
|
||||
cache_key, cached_data, kwargs = self.add_embedding_response_to_cache(
|
||||
result, kwargs["input"], kwargs
|
||||
)
|
||||
cache_list.append((cache_key, cached_data))
|
||||
elif isinstance(kwargs["input"], str):
|
||||
cache_key, cached_data, kwargs = self.add_embedding_response_to_cache(result, kwargs["input"], kwargs)
|
||||
cache_list.append((cache_key, cached_data))
|
||||
|
||||
if self._native_cache is not None:
|
||||
entries: Final = tuple(
|
||||
(request, cached_data["response"])
|
||||
for cache_key, cached_data in cache_list
|
||||
if (request := self._native_request(MappingProxyType({**kwargs, "cache_key": cache_key})))
|
||||
is not None
|
||||
)
|
||||
await self._native_cache.async_store_batch(
|
||||
tuple(request for request, _ in entries),
|
||||
tuple(response for _, response in entries),
|
||||
)
|
||||
elif dynamic_cache_object is not None:
|
||||
await dynamic_cache_object.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
else:
|
||||
await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
if self._native_cache is not None:
|
||||
entries: Final = tuple(
|
||||
(request, cached_data["response"])
|
||||
for cache_key, cached_data in cache_list
|
||||
if (request := self._native_request(MappingProxyType({**kwargs, "cache_key": cache_key})))
|
||||
is not None
|
||||
)
|
||||
await self._native_cache.async_store_batch(
|
||||
tuple(request for request, _ in entries),
|
||||
tuple(response for _, response in entries),
|
||||
)
|
||||
elif dynamic_cache_object is not None:
|
||||
await dynamic_cache_object.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
else:
|
||||
await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
except Exception as e:
|
||||
self._log_add_cache_failure(e)
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ import litellm
|
|||
from litellm._internal_context import post_response_phase
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.caching import InMemoryCache
|
||||
from litellm.caching.caching import S3Cache
|
||||
from litellm.caching.caching import S3Cache, response_cache_phase
|
||||
from litellm.constants import CACHE_WRITE_SHUTDOWN_FLUSH_TIMEOUT_SECONDS
|
||||
from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
|
||||
update_response_metadata,
|
||||
|
|
@ -146,16 +146,17 @@ _PENDING_CACHE_WRITES: Final[set["asyncio.Task[None]"]] = set() # mutable-ok: s
|
|||
|
||||
|
||||
async def _complete_cache_write_despite_cancellation(write_factory: Callable[[], Awaitable[None]]) -> None:
|
||||
try:
|
||||
await write_factory()
|
||||
except asyncio.CancelledError:
|
||||
with response_cache_phase("set"):
|
||||
try:
|
||||
await asyncio.wait_for(write_factory(), timeout=CACHE_WRITE_SHUTDOWN_FLUSH_TIMEOUT_SECONDS)
|
||||
except Exception as flush_error: # noqa: BLE001 # shutdown flush failures are logged, never raised
|
||||
verbose_logger.warning(
|
||||
"LiteLLM Cache: pending cache write failed during event loop shutdown: %s", flush_error
|
||||
)
|
||||
raise
|
||||
await write_factory()
|
||||
except asyncio.CancelledError:
|
||||
try:
|
||||
await asyncio.wait_for(write_factory(), timeout=CACHE_WRITE_SHUTDOWN_FLUSH_TIMEOUT_SECONDS)
|
||||
except Exception as flush_error: # noqa: BLE001 # shutdown flush failures are logged, never raised
|
||||
verbose_logger.warning(
|
||||
"LiteLLM Cache: pending cache write failed during event loop shutdown: %s", flush_error
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def create_cache_write_task(write_factory: Callable[[], Awaitable[None]]) -> "asyncio.Task[None]":
|
||||
|
|
@ -394,7 +395,8 @@ class LLMCachingHandler:
|
|||
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
|
||||
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
|
||||
print_verbose("Checking Sync Cache")
|
||||
cached_result = litellm.cache.get_cache(**new_kwargs)
|
||||
with response_cache_phase("get"):
|
||||
cached_result = litellm.cache.get_cache(**new_kwargs)
|
||||
if cached_result is not None:
|
||||
if "detail" in cached_result:
|
||||
# implies an error occurred
|
||||
|
|
@ -795,7 +797,7 @@ class LLMCachingHandler:
|
|||
new_kwargs["input"] = [new_kwargs["input"]]
|
||||
elif not isinstance(new_kwargs["input"], list):
|
||||
raise ValueError("input must be a string or a list")
|
||||
tasks: Final = []
|
||||
tasks: Final[list[Awaitable[object]]] = []
|
||||
for idx, i in enumerate(new_kwargs["input"]):
|
||||
preset_cache_key = litellm.cache.get_cache_key(**{**new_kwargs, "input": i})
|
||||
tasks.append(
|
||||
|
|
@ -804,7 +806,9 @@ class LLMCachingHandler:
|
|||
dynamic_cache_object=self.dual_cache,
|
||||
)
|
||||
)
|
||||
cached_result = [_current_format_embedding_entry(entry) for entry in await asyncio.gather(*tasks)]
|
||||
with response_cache_phase("get"):
|
||||
entries: Final = await asyncio.gather(*tasks)
|
||||
cached_result = [_current_format_embedding_entry(entry) for entry in entries]
|
||||
## check if cached result is None ##
|
||||
if cached_result is not None and isinstance(cached_result, list):
|
||||
# set cached_result to None if all elements are None
|
||||
|
|
@ -817,18 +821,20 @@ class LLMCachingHandler:
|
|||
if litellm.cache._supports_async() is True:
|
||||
## check if dual cache is supported ##
|
||||
self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
|
||||
cached_result = await litellm.cache.async_get_cache(
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
cache_key=self.preset_cache_key,
|
||||
**request_kwargs,
|
||||
)
|
||||
with response_cache_phase("get"):
|
||||
cached_result = await litellm.cache.async_get_cache(
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
cache_key=self.preset_cache_key,
|
||||
**request_kwargs,
|
||||
)
|
||||
else: # fallback for caches that don't support async
|
||||
self.preset_cache_key = request_cache_key or litellm.cache.get_cache_key(**request_kwargs)
|
||||
cached_result = litellm.cache.get_cache(
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
cache_key=self.preset_cache_key,
|
||||
**request_kwargs,
|
||||
)
|
||||
with response_cache_phase("get"):
|
||||
cached_result = litellm.cache.get_cache(
|
||||
dynamic_cache_object=self.dual_cache,
|
||||
cache_key=self.preset_cache_key,
|
||||
**request_kwargs,
|
||||
)
|
||||
return cached_result
|
||||
|
||||
def _convert_cached_result_to_model_response(
|
||||
|
|
@ -1118,7 +1124,8 @@ class LLMCachingHandler:
|
|||
return
|
||||
|
||||
if self._should_store_result_in_cache(original_function=self.original_function, kwargs=new_kwargs):
|
||||
litellm.cache.add_cache(result, **new_kwargs)
|
||||
with response_cache_phase("set"):
|
||||
litellm.cache.add_cache(result, **new_kwargs)
|
||||
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -22,9 +22,16 @@ from datetime import timedelta
|
|||
from types import MappingProxyType, TracebackType
|
||||
from typing import Final, Generic, Protocol, TypeVar
|
||||
|
||||
from litellm._internal_context import (
|
||||
REDIS_FAMILIES_METADATA_KEY,
|
||||
current_service_target,
|
||||
service_caller,
|
||||
service_target,
|
||||
)
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.redis_cache import (
|
||||
RedisCache,
|
||||
_get_call_stack_info, # pyright: ignore[reportPrivateUsage] # same caller chain every RedisCache method reports
|
||||
_run_under_circuit_breaker, # pyright: ignore[reportPrivateUsage] # same health signal as every RedisCache method
|
||||
log_redis_failure,
|
||||
)
|
||||
|
|
@ -56,12 +63,14 @@ class _Op(Generic[_T]):
|
|||
how to run on its own when the batch cannot pipeline (cluster client, or a reply the pipeline cannot
|
||||
settle, like NOSCRIPT)."""
|
||||
|
||||
__slots__ = ("future", "settled_hooks")
|
||||
__slots__ = ("caller", "future", "settled_hooks", "target")
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.future: Final[asyncio.Future[_T]] = asyncio.get_running_loop().create_future()
|
||||
self.future.add_done_callback(_mark_retrieved)
|
||||
self.settled_hooks: Final[list[SettledHook[_T]]] = [] # mutable-ok: append-only registry
|
||||
self.target: Final = current_service_target()
|
||||
self.caller: Final = _get_call_stack_info()
|
||||
|
||||
async def run_settled_hooks(self) -> None:
|
||||
for hook in self.settled_hooks:
|
||||
|
|
@ -100,7 +109,8 @@ class _Op(Generic[_T]):
|
|||
|
||||
async def _settle_alone(self) -> None:
|
||||
try:
|
||||
self.future.set_result(await self.run_alone())
|
||||
with service_target(self.target), service_caller(self.caller):
|
||||
self.future.set_result(await self.run_alone())
|
||||
except Exception as e: # noqa: BLE001 # the declaring caller owns the failure of its own operation
|
||||
self.future.set_exception(e)
|
||||
|
||||
|
|
@ -359,6 +369,7 @@ class RedisBatch:
|
|||
|
||||
async def _flush_pipeline(self, ops: Sequence[_Op[object]]) -> None:
|
||||
start_time: Final = time.time()
|
||||
target, metadata = _pipeline_service_event(ops)
|
||||
widths: list[int] = [] # mutable-ok: filled while enqueuing
|
||||
|
||||
async def run() -> list[object]:
|
||||
|
|
@ -371,28 +382,32 @@ class RedisBatch:
|
|||
replies: Final = await _run_under_circuit_breaker(self.redis_cache._circuit_breaker, self.name, run) # pyright: ignore[reportPrivateUsage] # same breaker as the cache's own methods
|
||||
except Exception as e: # noqa: BLE001 # each declaring caller applies its own Redis fallback
|
||||
log_redis_failure(verbose_logger, logging.WARNING, f"{self.name}: pipeline of {len(ops)} ops failed", e)
|
||||
asyncio.create_task(
|
||||
self.redis_cache.service_logger_obj.async_service_failure_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
error=e,
|
||||
call_type=f"{self.name}[{len(ops)}]",
|
||||
start_time=start_time,
|
||||
end_time=time.time(),
|
||||
with service_target(target):
|
||||
asyncio.create_task(
|
||||
self.redis_cache.service_logger_obj.async_service_failure_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
error=e,
|
||||
call_type=self.name,
|
||||
start_time=start_time,
|
||||
end_time=time.time(),
|
||||
event_metadata=metadata,
|
||||
)
|
||||
)
|
||||
)
|
||||
for op in ops:
|
||||
op.future.set_exception(e)
|
||||
return
|
||||
asyncio.create_task(
|
||||
self.redis_cache.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
call_type=f"{self.name}[{len(ops)}]",
|
||||
start_time=start_time,
|
||||
end_time=time.time(),
|
||||
with service_target(target):
|
||||
asyncio.create_task(
|
||||
self.redis_cache.service_logger_obj.async_service_success_hook(
|
||||
service=ServiceTypes.REDIS,
|
||||
duration=time.time() - start_time,
|
||||
call_type=self.name,
|
||||
start_time=start_time,
|
||||
end_time=time.time(),
|
||||
event_metadata=metadata,
|
||||
)
|
||||
)
|
||||
)
|
||||
retries: list[Awaitable[None]] = [] # mutable-ok: collected while slicing replies
|
||||
offset = 0
|
||||
for op, width in zip(ops, widths):
|
||||
|
|
@ -404,6 +419,18 @@ class RedisBatch:
|
|||
await asyncio.gather(*retries)
|
||||
|
||||
|
||||
MIXED_PIPELINE_TARGET: Final = "mixed"
|
||||
|
||||
|
||||
def _pipeline_service_event(ops: Sequence[_Op[object]]) -> tuple[str | None, dict[str, int | str]]:
|
||||
"""The target and metadata of one pipeline flush: the one key family every op was declared under, or
|
||||
``"mixed"`` plus the sorted families when owners of several families share the trip."""
|
||||
families: Final = sorted({op.target for op in ops if op.target is not None})
|
||||
if len(families) > 1:
|
||||
return MIXED_PIPELINE_TARGET, {"op_count": len(ops), REDIS_FAMILIES_METADATA_KEY: ",".join(families)}
|
||||
return next(iter(families), None), {"op_count": len(ops)}
|
||||
|
||||
|
||||
def _backend_key(redis_cache: RedisCache) -> object:
|
||||
"""Two ``RedisCache`` instances built from the same connection settings and namespace talk to the same server
|
||||
under the same key prefix, so the proxy's cache and the router's cache share one pipeline (the router gets its
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import asyncio
|
|||
import functools
|
||||
import hashlib
|
||||
import inspect
|
||||
import itertools
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
|
|
@ -21,12 +22,13 @@ from collections.abc import Awaitable, Callable, Iterator, Sequence
|
|||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
from types import MappingProxyType
|
||||
from types import FrameType, MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_service_caller
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import (
|
||||
DEFAULT_REDIS_MAJOR_VERSION,
|
||||
|
|
@ -92,9 +94,37 @@ class _AsyncRedisCommands(Protocol):
|
|||
def eval(self, script: str, numkeys: int, *keys_and_args: str | bytes | float) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
_BREAKER_GUARD_FRAME_NAMES: Final = frozenset(
|
||||
{"<lambda>", "wrapper", "_run_under_circuit_breaker", "_run_under_circuit_breaker_sync"}
|
||||
_GENERIC_CALLER_MODULES: Final = frozenset(
|
||||
{
|
||||
__name__,
|
||||
"litellm.caching.redis_batch",
|
||||
"litellm.caching.dual_cache",
|
||||
"litellm.caching.caching",
|
||||
"litellm.rust_bridge.lifecycle",
|
||||
"litellm.rust_bridge.streams",
|
||||
"contextlib",
|
||||
}
|
||||
)
|
||||
_GENERIC_CALLER_FRAME_NAMES: Final = frozenset(
|
||||
{
|
||||
"<lambda>",
|
||||
"wrapper",
|
||||
"_run_under_circuit_breaker",
|
||||
"_run_under_circuit_breaker_sync",
|
||||
"run_alone",
|
||||
"_settle_alone",
|
||||
"get_cache",
|
||||
"set_cache",
|
||||
"async_get_cache",
|
||||
"async_set_cache",
|
||||
"async_batch_get_cache",
|
||||
"async_batch_get_cache_shared",
|
||||
"async_set_cache_pipeline",
|
||||
"async_increment_cache",
|
||||
"async_delete_cache",
|
||||
}
|
||||
)
|
||||
_CALL_STACK_END_MODULES: Final = ("asyncio", "concurrent", "threading")
|
||||
|
||||
_INCREMENT_WITH_FLOOR_LUA: Final = (
|
||||
"local count = redis.call('INCRBY', KEYS[1], ARGV[1]) "
|
||||
|
|
@ -113,18 +143,39 @@ def _decoded_counts(values: Sequence[bytes | str | None]) -> tuple[int | None, .
|
|||
)
|
||||
|
||||
|
||||
def _is_generic_caller_frame(frame: FrameType) -> bool:
|
||||
module: Final = frame.f_globals.get("__name__")
|
||||
return module in _GENERIC_CALLER_MODULES or frame.f_code.co_name in _GENERIC_CALLER_FRAME_NAMES
|
||||
|
||||
|
||||
def _ends_call_stack(frame: FrameType) -> bool:
|
||||
module: Final = frame.f_globals.get("__name__")
|
||||
return isinstance(module, str) and module.startswith(_CALL_STACK_END_MODULES)
|
||||
|
||||
|
||||
def _caller_frames(first: FrameType) -> Iterator[FrameType]:
|
||||
frame: FrameType | None = first
|
||||
while frame is not None and not _ends_call_stack(frame):
|
||||
yield frame
|
||||
frame = frame.f_back
|
||||
|
||||
|
||||
def _get_call_stack_info(num_frames: int = 2) -> str:
|
||||
"""
|
||||
Get the function names from the previous 1-2 functions in the call stack.
|
||||
Get the function names of the nearest meaningful callers of the cache method.
|
||||
|
||||
Frames belonging to this module's circuit-breaker guards are skipped so the
|
||||
reported callers stay the real ones even on guarded methods.
|
||||
Frames that merely forward the call (this module's circuit-breaker guards, the
|
||||
cache facades, the batch pipeline's retry path, generic cache verbs) are
|
||||
skipped, and the walk stops at the event loop, so the chain names the litellm
|
||||
code that wanted the call. When nothing but forwarding frames is found (the call
|
||||
runs in a task of its own, like a batch op retried on the flush) the chain the
|
||||
declaring code threaded through ``service_caller`` is reported, else ``unknown``.
|
||||
|
||||
Args:
|
||||
num_frames: Number of previous frames to include (default: 2)
|
||||
|
||||
Returns:
|
||||
A string with format "current_function <- caller_function [<- grandparent_function]"
|
||||
A string with format "caller_function [<- grandparent_function]"
|
||||
"""
|
||||
try:
|
||||
current_frame: Final = inspect.currentframe()
|
||||
|
|
@ -135,22 +186,23 @@ def _get_call_stack_info(num_frames: int = 2) -> str:
|
|||
f_back: Final = current_frame.f_back
|
||||
if f_back is None:
|
||||
return "unknown"
|
||||
frame = f_back.f_back
|
||||
if frame is None:
|
||||
first: Final = f_back.f_back
|
||||
if first is None:
|
||||
return "unknown"
|
||||
function_names: Final = []
|
||||
frames: Final = _caller_frames(first)
|
||||
leading: Final = tuple(itertools.islice(frames, num_frames))
|
||||
leading_names: Final = tuple(frame.f_code.co_name for frame in leading if not _is_generic_caller_frame(frame))
|
||||
further_names: Final = tuple(
|
||||
itertools.islice(
|
||||
(frame.f_code.co_name for frame in frames if not _is_generic_caller_frame(frame)),
|
||||
num_frames - len(leading_names),
|
||||
)
|
||||
)
|
||||
function_names: Final = leading_names + further_names
|
||||
|
||||
while frame is not None and len(function_names) < num_frames:
|
||||
if frame.f_code.co_name in _BREAKER_GUARD_FRAME_NAMES and frame.f_globals.get("__name__") == __name__:
|
||||
frame = frame.f_back
|
||||
continue
|
||||
function_names.append(frame.f_code.co_name)
|
||||
frame = frame.f_back
|
||||
|
||||
if not function_names:
|
||||
return "unknown"
|
||||
|
||||
return " <- ".join(function_names)
|
||||
if function_names:
|
||||
return " <- ".join(function_names)
|
||||
return current_service_caller() or "unknown"
|
||||
except Exception:
|
||||
return "unknown"
|
||||
|
||||
|
|
@ -1141,7 +1193,6 @@ class RedisCache(BaseCache):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
event_metadata={"key": key},
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
|
@ -1158,7 +1209,6 @@ class RedisCache(BaseCache):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
event_metadata={"key": key},
|
||||
)
|
||||
)
|
||||
log_redis_failure(
|
||||
|
|
@ -1669,7 +1719,6 @@ class RedisCache(BaseCache):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
event_metadata={"key": key},
|
||||
)
|
||||
)
|
||||
return response
|
||||
|
|
@ -1686,7 +1735,6 @@ class RedisCache(BaseCache):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
event_metadata={"key": key},
|
||||
)
|
||||
)
|
||||
print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {e}")
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import time
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
|
|
@ -21,6 +22,8 @@ from litellm.types.integrations.slack_alerting import (
|
|||
HangingRequestData,
|
||||
)
|
||||
|
||||
_REQUEST_STATUS_TARGET: Final = "request_status"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
else:
|
||||
|
|
@ -82,6 +85,7 @@ class AlertingHangingRequestCheck:
|
|||
)
|
||||
return
|
||||
|
||||
@with_service_target(_REQUEST_STATUS_TARGET)
|
||||
async def send_alerts_for_hanging_requests(self):
|
||||
"""
|
||||
Send alerts for hanging requests
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ import litellm
|
|||
import litellm.litellm_core_utils
|
||||
import litellm.litellm_core_utils.litellm_logging
|
||||
import litellm.types
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -83,6 +84,9 @@ def _proxy_llm_router() -> Router | None:
|
|||
return llm_router
|
||||
|
||||
|
||||
_DAILY_REPORT_TARGET: Final = "daily_report_schedule"
|
||||
|
||||
|
||||
class SlackAlerting(CustomBatchLogger):
|
||||
"""
|
||||
Class for sending Slack Alerts
|
||||
|
|
@ -1760,18 +1764,20 @@ Model Info:
|
|||
"""
|
||||
report_sent_bool = False
|
||||
|
||||
report_sent: Final = await self.internal_usage_cache.async_get_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
parent_otel_span=None,
|
||||
) # None | float
|
||||
with service_target(_DAILY_REPORT_TARGET):
|
||||
report_sent: Final = await self.internal_usage_cache.async_get_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
parent_otel_span=None,
|
||||
) # None | float
|
||||
|
||||
current_time: Final = time.time()
|
||||
|
||||
if report_sent is None:
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
value=current_time,
|
||||
)
|
||||
with service_target(_DAILY_REPORT_TARGET):
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
value=current_time,
|
||||
)
|
||||
elif isinstance(report_sent, float):
|
||||
# Check if current time - interval >= time last sent
|
||||
interval_seconds: Final = self.alerting_args.daily_report_frequency
|
||||
|
|
@ -1790,10 +1796,11 @@ Model Info:
|
|||
# Sneak in the reporting logic here
|
||||
await self.send_daily_reports(router=llm_router)
|
||||
# Also, don't forget to update the report_sent time after sending the report!
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
value=current_time,
|
||||
)
|
||||
with service_target(_DAILY_REPORT_TARGET):
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=SlackAlertingCacheKeys.report_sent_key.value,
|
||||
value=current_time,
|
||||
)
|
||||
report_sent_bool = True
|
||||
|
||||
return report_sent_bool
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_a
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -55,6 +56,8 @@ from litellm.exceptions import (
|
|||
SensitiveDataRouteException,
|
||||
)
|
||||
|
||||
GUARDRAIL_SESSIONS_TARGET: Final = "guardrail_sessions"
|
||||
|
||||
# Per-process secret tagging each recorded marker. The deployment hook only
|
||||
# honors markers carrying this token, so a caller cannot forge the metadata
|
||||
# field to suppress a guardrail on the direct-SDK path that never reaches the
|
||||
|
|
@ -474,6 +477,7 @@ class CustomGuardrail(CustomLogger):
|
|||
def _scanned_texts_cache_key(self, session_id: str) -> str:
|
||||
return f"guardrail_scanned_texts:{self.guardrail_name}:{session_id}"
|
||||
|
||||
@with_service_target(GUARDRAIL_SESSIONS_TARGET)
|
||||
async def filter_new_texts_for_session(
|
||||
self,
|
||||
texts: list[str] | None,
|
||||
|
|
@ -518,6 +522,7 @@ class CustomGuardrail(CustomLogger):
|
|||
seen: Final[set[str]] = {str(h) for h in cached} if isinstance(cached, list) else set()
|
||||
return [text for text in texts if self._scanned_text_hash(text) not in seen]
|
||||
|
||||
@with_service_target(GUARDRAIL_SESSIONS_TARGET)
|
||||
async def mark_texts_scanned(
|
||||
self,
|
||||
texts: list[str] | None,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,10 @@ SERVER span "POST /v1/chat/completions" ← FastAPI instrumentation
|
|||
│ ├── CLIENT span "postgres get_key_object" ← datastore call │
|
||||
│ └── CLIENT span "postgres get_team_membership" │
|
||||
├── INTERNAL span "execute_guardrail …" ← guardrail │ this package
|
||||
├── INTERNAL span "cache.get llm_response" ← response cache │
|
||||
│ └── CLIENT span "redis.get llm_response" │
|
||||
├── INTERNAL span "route gpt-4o" ← deployment pick │
|
||||
│ └── CLIENT span "redis.mget router_cooldowns" │
|
||||
├── CLIENT span "chat gpt-4o" ← LLM call │
|
||||
└── CLIENT span "batch_write_to_db …" ← spend write ┘
|
||||
```
|
||||
|
|
@ -59,11 +63,69 @@ traceable units of work:
|
|||
the trace. `auth` is also excluded here because it gets a **live phase span**
|
||||
instead (see below).
|
||||
|
||||
Spans are named `"{service} {call_type}"` (e.g. `"redis set"`) so repeated calls
|
||||
to one service stay distinguishable. `call_type` is the operation only; the
|
||||
litellm call chain that issued it (`async_set_cache <- async_add_cache`) travels
|
||||
as `ServiceLoggerPayload.caller` and lands on the `litellm.service.caller`
|
||||
attribute, so one operation is one span name. Like every other span they parent to the
|
||||
Redis spans are named `"{service}.{verb} {target}"` (e.g. `"redis.get llm_response"`,
|
||||
`"redis.mget auth_objects"`), the `{db.operation.name} {target}` shape of the OTel
|
||||
database conventions: the verb comes from the cache method
|
||||
(`spans._SERVICE_VERB_BY_CALL_TYPE`), the target from the producer running the
|
||||
call inside `litellm._internal_context.service_target(...)` and is a key family
|
||||
(`llm_response`, `auth_objects`, `router_cooldowns`, `router_cooldowns_usage`,
|
||||
`router_usage`, `router_budgets`, `router_session_pins`, `rate_limits`,
|
||||
`model_budgets`, `session_budgets`, `session_iterations`, `sensitive_route_pins`,
|
||||
`prompt_cache_pins`, `prompt_cache_predictions`, `spend_counters`, `config_params`,
|
||||
`daily_report_schedule`), never a key. The whole `auth` phase runs under
|
||||
`auth_objects`, so every cache read it triggers is `redis.get auth_objects` /
|
||||
`redis.mget auth_objects`, and so does the post-call spend write-back into the
|
||||
same auth objects. A proxy hook or routing strategy declares its family once, on
|
||||
its entrypoints, with `@with_service_target("rate_limits")`, so every read and
|
||||
write it issues (helpers included) carries it; the response-cache facade
|
||||
(`Cache.get_cache` / `async_get_cache` / `add_cache` / `async_add_cache` /
|
||||
`async_add_cache_pipeline`) opens the `cache.get llm_response` /
|
||||
`cache.set llm_response` phase itself, so a lookup issued by the native bridge
|
||||
is phased and targeted like one issued by `caching_handler.py`. The verb is the
|
||||
Redis command the method issues (`get`, `mget`, `set`, `sadd`, `incr`, `ttl`,
|
||||
`expire`, `delete`, `rpush`, `lpop`, `scan`, `ping`), so the cooldown fail counter
|
||||
shows as `redis.incr router_cooldowns` followed by `redis.ttl router_cooldowns` /
|
||||
`redis.expire router_cooldowns`. Background producers declare a family the same
|
||||
way (`pod_lock`, `budget_reset`, `spend_queue`, `health_check`, `scheduler_queue`,
|
||||
`managed_files`, `mcp_servers`, ...), so a job tick renders `redis.set pod_lock`
|
||||
rather than a bare `redis.set`; `tests/unit/test_internal_context.py` scans every
|
||||
module under `litellm/` and `enterprise/` that calls a shared cache or declares a
|
||||
read or write on the request batch (`reserve_redis_batch_reads`,
|
||||
`declare_batch_get`, `batch.mget`, `batch.set`, `batch.script`) and fails when
|
||||
one has no declared family, with the process-local `InMemoryCache` callers listed
|
||||
as the only exemptions. A batch op carries the family that was active when it was
|
||||
declared, so the routing prefetch armed before deployment selection
|
||||
(`RoutingPrefetch.arm`) is `router_cooldowns` when only cooldown keys go out,
|
||||
`router_usage` when only usage counters do, and `router_cooldowns_usage` when both
|
||||
ride the same MGET, whichever pipeline or standalone read later settles it. A
|
||||
per-request pipeline (`RedisBatch`) that carries ops of
|
||||
one family is `"redis.pipeline auth_objects"`; one that carries several owners'
|
||||
ops is `"redis.pipeline mixed"` with the sorted family list on
|
||||
`litellm.redis.families` and the op count on `litellm.metadata.op_count` (an int,
|
||||
never stringified). A cluster client cannot pipeline across slots, so there every
|
||||
batch op settles on its own and one write-back of three auth objects shows as three
|
||||
parallel `redis.set auth_objects` spans with the same caller, not one pipeline
|
||||
span. Every `call_type` the Redis cache layer emits maps to a verb,
|
||||
so the `{service} {call_type}` fallback is unreachable for Redis (a test asserts
|
||||
it). Postgres spans are `postgres.{verb} {table}`, see
|
||||
https://github.com/BerriAI/litellm/pull/44240; the other non-Redis services keep
|
||||
the `"{service} {call_type}"` name (`"batch_write_to_db _PROXY_track_cost_callback"`):
|
||||
one scheme, `{service}.{verb} {target}` when the method maps to a verb and
|
||||
`{service} {call_type}` otherwise, and never a count, key or id in the name. Either way
|
||||
the raw method name stays on `litellm.service.call_type` and `db.operation.name`
|
||||
(and the bare `call_type` the metrics are keyed by), the target lands on
|
||||
`litellm.service.target`, and the litellm call chain that issued the call
|
||||
(`_retrieve_from_cache <- _async_get_cache`) travels as
|
||||
`ServiceLoggerPayload.caller` onto `litellm.service.caller`, with the forwarding
|
||||
frames (cache facades, circuit-breaker guards, batch retry wrappers, the native
|
||||
execution's `lifecycle`/`streams` drivers) skipped so it names the code that wanted
|
||||
the call. A call whose own frames are all forwarders
|
||||
(a batch op settled in a task of its own, on a cluster client or a NOSCRIPT retry)
|
||||
reports the chain its declaring code captured and threaded through
|
||||
`service_caller(...)`, never the forwarders, and `unknown` when there is none.
|
||||
The cache key itself is never on the span: it is unbounded and carries key hashes
|
||||
and session ids, and the span is already named by key family. Like every other
|
||||
span they parent to the
|
||||
**ambient** context, falling back to the threaded `litellm_parent_otel_span` only
|
||||
when ambient has no live span; a background job with neither starts its own root
|
||||
trace.
|
||||
|
|
@ -94,7 +156,21 @@ Caller-supplied `event_metadata` is **sanitized** before it reaches a span
|
|||
|
||||
**Live phase spans.** `auth` is wrapped in a real, active span
|
||||
(`logger.phase_span`) for the duration of authentication, so the DB lookups it
|
||||
triggers nest **under** it instead of flattening onto the server span. Identity
|
||||
triggers nest **under** it instead of flattening onto the server span. The
|
||||
response cache does the same: the lookup runs inside `cache.get llm_response`
|
||||
(a child of the server span, so its Redis read sits before `chat {model}` in
|
||||
causal order) and the write inside `cache.set llm_response`. Deployment selection
|
||||
runs inside `route {model_group}` (`Router.async_get_available_deployment`, the
|
||||
requested group, never the deployment it picks), so the cooldown, usage and
|
||||
model-id reads the router issues nest under it, before `chat {model}`; the phase
|
||||
is opened in Python, never inside the native lifecycle. The one known ordering
|
||||
limitation is the native path (`LITELLM_RUST`): its lifecycle fires pre-call
|
||||
logging before it yields the cache await, so `cache.get llm_response` starts after
|
||||
`chat {model}` there, and moving it needs native changes that
|
||||
`litellm/rust_bridge/AGENTS.md` forbids. The write runs from
|
||||
the post-response phase, so that span is a linked root rather than a child that
|
||||
would stretch the request, and the Redis write it issues nests under it instead
|
||||
of starting a third trace (`context.post_response_root`). Identity
|
||||
Baggage (team/key/user) is seeded once the key resolves, so every post-auth span
|
||||
inherits it; auth-internal DB lookups that run before the key is known stay
|
||||
unlabeled, which is correct.
|
||||
|
|
@ -163,7 +239,9 @@ becomes the global, so server spans export to that backend too.
|
|||
sync-only provider driven through a thread pool, where contextvars (and so the
|
||||
anchor) don't follow — no parent is visible there, so creation is **deferred**
|
||||
to the async callback, whose worker context was copied from the request task at
|
||||
enqueue and so still carries the anchor. **Pass-through** endpoints call
|
||||
enqueue and so still carries the anchor. A deferred span starts at the provider
|
||||
handoff (`api_call_start_time`), not at the logging object's creation, so it
|
||||
bounds the provider attempt rather than the whole request. **Pass-through** endpoints call
|
||||
`logging_obj.pre_call` in the request task too, then close from a detached
|
||||
`asyncio.create_task`; the anchor (not the by-then-inactive server span) keeps
|
||||
their LLM-call span in the request's trace. `pre_call` is litellm's generic
|
||||
|
|
@ -326,7 +404,13 @@ lives in [`plumbing/`](./plumbing):
|
|||
(`DYNAMIC_HEADERS_BY_CALLBACK`). Presets do **no** network I/O at build time:
|
||||
AgentOps, for example, mints its JWT lazily inside a custom exporter on the
|
||||
first export (in the `BatchSpanProcessor` worker thread), never on the event
|
||||
loop.
|
||||
loop. A preset built while another `OpenTelemetryV2` logger is already
|
||||
registered (a key or team `logging` entry naming `arize`, say, beside the
|
||||
operator's `otel`) keeps only the exporters it contributed itself: the
|
||||
registered logger already delivers every call to the operator's collector, so
|
||||
a copy of those base exporters would emit each `chat` span there twice. A
|
||||
preset that contributes no exporter of its own (Langtrace is a mapper over the
|
||||
operator's collector) keeps the base exporters it has nothing to replace with.
|
||||
|
||||
## Extending
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from dataclasses import replace
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
|
|
@ -52,10 +52,13 @@ from litellm.integrations.otel.model.semconv import Error
|
|||
from litellm.integrations.otel.model.spans import SpanRole, span_role_for_service
|
||||
from litellm.integrations.otel.model.utils import to_ns
|
||||
from litellm.integrations.otel.plumbing.context import (
|
||||
active_phase,
|
||||
is_recordable_span,
|
||||
mcp_message_transport_span,
|
||||
post_response_root,
|
||||
request_root_http_route,
|
||||
request_root_span,
|
||||
resolve_internal_call_span_context,
|
||||
resolve_mcp_span_context,
|
||||
resolve_request_span_context,
|
||||
resolve_service_span_context,
|
||||
|
|
@ -175,6 +178,12 @@ class _LLMCallSpan:
|
|||
self.provider = provider
|
||||
|
||||
|
||||
def _llm_call_parent_context(call: LLMCallEvent) -> Context:
|
||||
"""A call litellm makes on the request's behalf (a classifier, a judge) parents under the
|
||||
phase that made it; the provider attempt parents under the request root."""
|
||||
return resolve_internal_call_span_context() if call.purpose is not None else resolve_request_span_context()
|
||||
|
||||
|
||||
class OpenTelemetryV2(CustomLogger):
|
||||
"""The ``CustomLogger`` for OpenTelemetry."""
|
||||
|
||||
|
|
@ -310,7 +319,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
# callback (the thread-pool case, where the anchor isn't visible here).
|
||||
# Do not route on the deferred path: creating or LRU-touching a tenant
|
||||
# provider here would evict idle ones even though close re-routes.
|
||||
parent_context: Final = resolve_request_span_context()
|
||||
parent_context: Final = _llm_call_parent_context(call)
|
||||
if not is_recordable_span(get_current_span(parent_context)):
|
||||
self._store_open_call(call_id, _LLMCallSpan(span=None, start_time_ns=start_time_ns))
|
||||
return
|
||||
|
|
@ -561,6 +570,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
capture_content=self.config.capture_span_content,
|
||||
time_to_first_chunk_seconds=call.time_to_first_chunk_seconds,
|
||||
request_route=request_root_http_route(),
|
||||
request_purpose=call.purpose,
|
||||
trace=call.trace,
|
||||
session_id=call.session_id,
|
||||
)
|
||||
|
|
@ -579,16 +589,22 @@ class OpenTelemetryV2(CustomLogger):
|
|||
# root span — parent to it (ambient fallback on the SDK path). Seed identity
|
||||
# Baggage so the span — and the SDK path, which has none — is labeled
|
||||
# consistently. A detached route roots its own trace instead, linked back.
|
||||
# With no carrier the span starts at the provider handoff, so a destination
|
||||
# logger's copy bounds the provider attempt like the operator's does.
|
||||
route: Final = self._tenant_tracers.route_for(self.tracer, call.dynamic_params, call.auth_metadata)
|
||||
try:
|
||||
parent_ctx: Final = self._seed_identity_baggage(
|
||||
data.identity, data.request_model, resolve_request_span_context()
|
||||
data.identity, data.request_model, _llm_call_parent_context(call)
|
||||
)
|
||||
return self._emitter.emit(
|
||||
SpanRole.LLM_CALL,
|
||||
data,
|
||||
parent_context=(set_span_in_context(INVALID_SPAN, parent_ctx) if route.detached else parent_ctx),
|
||||
start_time_ns=(carrier.start_time_ns if carrier is not None else to_ns(start_time)),
|
||||
start_time_ns=(
|
||||
carrier.start_time_ns
|
||||
if carrier is not None
|
||||
else to_ns(call.upstream_start_seconds) or to_ns(start_time)
|
||||
),
|
||||
end_time_ns=end_time_ns,
|
||||
tracer=route.tracer,
|
||||
links=_request_trace_links(parent_ctx) if route.detached else None,
|
||||
|
|
@ -735,8 +751,19 @@ class OpenTelemetryV2(CustomLogger):
|
|||
|
||||
@contextmanager
|
||||
def start_phase_span(self, name: str) -> "Iterator[Span]":
|
||||
span: Final = self._emitter.start_span(SpanRole.SERVICE, name)
|
||||
with use_span(span, end_on_exit=True):
|
||||
"""A live INTERNAL span the service calls inside the block nest under.
|
||||
|
||||
Parents like a service span: ambient first, and from the post-response phase
|
||||
it becomes a linked root that then adopts the calls made inside it, so the
|
||||
response-cache write is one small trace rather than a scatter of roots.
|
||||
"""
|
||||
parent_context, links = resolve_service_span_context()
|
||||
span: Final = self._emitter.start_span(SpanRole.SERVICE, name, parent_context=parent_context, links=links)
|
||||
with (
|
||||
use_span(span, end_on_exit=True),
|
||||
active_phase(span),
|
||||
post_response_root(span) if links else nullcontext(),
|
||||
):
|
||||
try:
|
||||
yield span
|
||||
except Exception as exc:
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ table: one lambda per mapping operation, applied against the typed span data.
|
|||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
from litellm._internal_context import REDIS_FAMILIES_METADATA_KEY
|
||||
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData
|
||||
from litellm.integrations.otel.mappers.utils import (
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
|
||||
|
|
@ -91,6 +92,7 @@ class GenAIMapper:
|
|||
f"{LiteLLM.COST_PREFIX}margin_total_amount": lambda d: d.cost.margin_total_amount,
|
||||
LiteLLM.REQUEST_STREAMING: lambda d: d.is_streaming,
|
||||
LiteLLM.REQUEST_ROUTE: lambda d: d.request_route,
|
||||
LiteLLM.REQUEST_PURPOSE: lambda d: d.request_purpose,
|
||||
}
|
||||
|
||||
_TOOL_ATTRS: dict[str, Callable[[ToolDefinition], AttrValue | None]] = {
|
||||
|
|
@ -149,6 +151,7 @@ class GenAIMapper:
|
|||
LiteLLM.SERVICE_NAME: lambda d: d.service_name,
|
||||
LiteLLM.SERVICE_CALL_TYPE: lambda d: d.call_type,
|
||||
LiteLLM.SERVICE_CALLER: lambda d: d.caller,
|
||||
LiteLLM.SERVICE_TARGET: lambda d: d.target,
|
||||
}
|
||||
|
||||
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
|
||||
|
|
@ -194,5 +197,12 @@ class GenAIMapper:
|
|||
# semconv naming the server it reached. Internal services (router, budget
|
||||
# jobs, …) have no db.system, so they get only the litellm.service.* keys.
|
||||
attrs.update(db_span_attributes(data.service_name, data.call_type))
|
||||
attrs.update({f"{LiteLLM.METADATA_PREFIX}{key}": value for key, value in data.event_metadata.items()})
|
||||
attrs.update(
|
||||
{
|
||||
LiteLLM.REDIS_FAMILIES
|
||||
if key == REDIS_FAMILIES_METADATA_KEY
|
||||
else f"{LiteLLM.METADATA_PREFIX}{key}": value
|
||||
for key, value in data.event_metadata.items()
|
||||
}
|
||||
)
|
||||
return attrs
|
||||
|
|
|
|||
|
|
@ -38,17 +38,25 @@ from __future__ import annotations
|
|||
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, cast, get_args
|
||||
|
||||
from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL, SESSION_ID_GENERATED_METADATA_KEY
|
||||
from litellm.constants import (
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL,
|
||||
SESSION_ID_GENERATED_METADATA_KEY,
|
||||
)
|
||||
from litellm.integrations.otel.model.semconv import resolve_operation
|
||||
from litellm.integrations.otel.model.trace_controls import TraceControls, caller_trace_controls
|
||||
from litellm.integrations.otel.model.utils import as_str, as_str_mapping, to_seconds
|
||||
from litellm.types.utils import InternalCallOrigin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
_INTERNAL_CALL_ORIGINS: Final[frozenset[str]] = frozenset(get_args(InternalCallOrigin))
|
||||
|
||||
REQUESTER_METADATA_KEY: Final = "requester_metadata"
|
||||
REQUESTER_METADATA_PATH: Final = f"{REQUESTER_METADATA_KEY}."
|
||||
|
||||
|
|
@ -220,6 +228,14 @@ class LLMCallEvent:
|
|||
# actually attempted — router pre-call rejections, SDK failures before the
|
||||
# provider handoff, and standalone guardrail runs all lack it.
|
||||
upstream_started: bool
|
||||
# When the request handed off to the provider, in epoch seconds. A close with no
|
||||
# carrier (a destination logger never sees ``pre_call``) starts its span here,
|
||||
# not at the logging object's creation, which predates routing and the cache.
|
||||
upstream_start_seconds: float | None
|
||||
# The litellm feature that made this call on the caller's behalf (an
|
||||
# ``InternalCallOrigin`` such as ``autorouter_classifier``), ``None`` for the
|
||||
# caller's own provider attempt.
|
||||
purpose: str | None
|
||||
# A best-effort ``"{operation} {model}"`` name known at ``pre_call`` time. The
|
||||
# span is renamed from the typed payload at close (``finish_span``); this only
|
||||
# needs to be reasonable for a span that never gets closed (a leak).
|
||||
|
|
@ -242,6 +258,8 @@ class LLMCallEvent:
|
|||
auth_metadata=auth_metadata(payload, kwargs),
|
||||
is_no_upstream_call=bool(kwargs.get(LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL)),
|
||||
upstream_started=kwargs.get("api_call_start_time") is not None,
|
||||
upstream_start_seconds=_epoch_seconds(kwargs.get("api_call_start_time")),
|
||||
purpose=internal_call_origin(payload, kwargs),
|
||||
provisional_span_name=f"{operation.value} {model}".strip(),
|
||||
time_to_first_chunk_seconds=time_to_first_chunk_seconds(kwargs),
|
||||
trace=trace,
|
||||
|
|
@ -249,6 +267,22 @@ class LLMCallEvent:
|
|||
)
|
||||
|
||||
|
||||
def _epoch_seconds(value: object) -> float | None:
|
||||
return to_seconds(value) if isinstance(value, (datetime, float, int, str)) and not isinstance(value, bool) else None
|
||||
|
||||
|
||||
def internal_call_origin(payload: StandardLoggingPayload | None, kwargs: Mapping[str, object]) -> str | None:
|
||||
"""The ``InternalCallOrigin`` a litellm-made sub-call carries in its request metadata, else ``None``."""
|
||||
return next(
|
||||
(
|
||||
origin
|
||||
for metadata in _metadata_dicts(payload, kwargs)
|
||||
if (origin := as_str(metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY))) in _INTERNAL_CALL_ORIGINS
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def caller_session_id(kwargs: Mapping[str, object], trace: TraceControls) -> str | None:
|
||||
"""The conversation id the caller sent (``litellm_session_id``, else the
|
||||
``session_id`` trace control); ``None`` when the request carried none.
|
||||
|
|
|
|||
|
|
@ -324,17 +324,21 @@ class GuardrailSpanData:
|
|||
)
|
||||
|
||||
|
||||
MetadataScalar = str | int | float | bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ServiceSpanData:
|
||||
service_name: str
|
||||
call_type: str | None = None
|
||||
caller: str | None = None
|
||||
target: str | None = None
|
||||
error: SpanError | None = None
|
||||
# Caller-supplied attributes to stamp on the service span, passed through
|
||||
# from ``async_service_*_hook(event_metadata=...)``. The mapper owns how
|
||||
# these are namespaced: the canonical vocabulary uses ``litellm.metadata.*``
|
||||
# keys, the semconv-ai / Traceloop vocabulary uses the bare key names.
|
||||
event_metadata: Mapping[str, str] = field(default_factory=dict)
|
||||
event_metadata: Mapping[str, MetadataScalar] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_payload(
|
||||
|
|
@ -351,6 +355,7 @@ class ServiceSpanData:
|
|||
service_name=payload.service.value,
|
||||
call_type=payload.call_type,
|
||||
caller=payload.caller,
|
||||
target=payload.target,
|
||||
error=SpanError(message=payload.error) if payload.error else None,
|
||||
event_metadata=sanitize_event_metadata(event_metadata),
|
||||
)
|
||||
|
|
@ -427,6 +432,7 @@ class LLMCallSpanData:
|
|||
output_type: GenAIOutputType | None = None
|
||||
call_type: str | None = None
|
||||
request_route: str | None = None
|
||||
request_purpose: str | None = None
|
||||
trace: TraceControls = field(default_factory=TraceControls)
|
||||
session_id: str | None = None
|
||||
embedding_output: EmbeddingOutput | None = None
|
||||
|
|
@ -438,6 +444,7 @@ class LLMCallSpanData:
|
|||
capture_content: bool = False,
|
||||
time_to_first_chunk_seconds: float | None = None,
|
||||
request_route: str | None = None,
|
||||
request_purpose: str | None = None,
|
||||
trace: TraceControls | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> LLMCallSpanData:
|
||||
|
|
@ -485,6 +492,7 @@ class LLMCallSpanData:
|
|||
output_type=resolve_output_type(call_type),
|
||||
call_type=call_type or None,
|
||||
request_route=request_route or context.identity.request_route,
|
||||
request_purpose=request_purpose,
|
||||
trace=trace or TraceControls(),
|
||||
session_id=session_id or None,
|
||||
embedding_output=embedding_output if capture_content else None,
|
||||
|
|
@ -649,17 +657,18 @@ _MAX_METADATA_ITEMS: Final = 32
|
|||
|
||||
def sanitize_event_metadata(
|
||||
event_metadata: Mapping[str, object] | None,
|
||||
) -> dict[str, str]:
|
||||
"""Reduce caller-supplied ``event_metadata`` to span-safe string attributes.
|
||||
) -> dict[str, MetadataScalar]:
|
||||
"""Reduce caller-supplied ``event_metadata`` to span-safe primitive attributes.
|
||||
|
||||
Keeps only primitive values (str/int/float/bool) under non-sensitive keys —
|
||||
never ``repr()``-ing objects, dicts, or lists, never stamping secrets/headers,
|
||||
and bounding the count and per-value length. This is the single chokepoint:
|
||||
both the GenAI and legacy mappers read the cleaned result.
|
||||
Keeps only primitive values (str/int/float/bool, each in its own type so a
|
||||
count stays a number) under non-sensitive keys — never ``repr()``-ing objects,
|
||||
dicts, or lists, never stamping secrets/headers, and bounding the count and
|
||||
per-string length. This is the single chokepoint: both the GenAI and legacy
|
||||
mappers read the cleaned result.
|
||||
"""
|
||||
if not event_metadata:
|
||||
return {}
|
||||
clean: Final[dict[str, str]] = {}
|
||||
clean: Final[dict[str, MetadataScalar]] = {}
|
||||
for key, value in event_metadata.items():
|
||||
if len(clean) >= _MAX_METADATA_ITEMS:
|
||||
break
|
||||
|
|
@ -670,8 +679,10 @@ def sanitize_event_metadata(
|
|||
continue
|
||||
# ``bool`` is a subclass of ``int``, so it's covered. Non-primitive values
|
||||
# (objects, dicts, lists) are dropped rather than stringified.
|
||||
if isinstance(value, (str, int, float)):
|
||||
clean[key] = str(value)[:_MAX_METADATA_VALUE_LEN]
|
||||
if isinstance(value, str):
|
||||
clean[key] = value[:_MAX_METADATA_VALUE_LEN]
|
||||
elif isinstance(value, (int, float)):
|
||||
clean[key] = value
|
||||
return clean
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -300,6 +300,9 @@ class LiteLLM:
|
|||
PROVIDER_MODEL: Final = "litellm.provider.model"
|
||||
REQUEST_STREAMING: Final = "litellm.request.streaming"
|
||||
REQUEST_ROUTE: Final = "litellm.request.route"
|
||||
# Which litellm feature made this LLM call when it is not the caller's own
|
||||
# provider attempt (e.g. ``autorouter_classifier``); absent on the real call.
|
||||
REQUEST_PURPOSE: Final = "litellm.request.purpose"
|
||||
TOOLS_DECLARED: Final = "litellm.request.tools.declared"
|
||||
GUARDRAIL_NAME: Final = "litellm.guardrail.name"
|
||||
GUARDRAIL_MODE: Final = "litellm.guardrail.mode"
|
||||
|
|
@ -327,6 +330,9 @@ class LiteLLM:
|
|||
SERVICE_NAME: Final = "litellm.service.name"
|
||||
SERVICE_CALL_TYPE: Final = "litellm.service.call_type"
|
||||
SERVICE_CALLER: Final = "litellm.service.caller"
|
||||
SERVICE_TARGET: Final = "litellm.service.target"
|
||||
# The sorted, comma-joined key families one Redis pipeline carried ops for; bounded, unlike the keys.
|
||||
REDIS_FAMILIES: Final = "litellm.redis.families"
|
||||
PREPROCESSING_MS: Final = "litellm.preprocessing.duration_ms"
|
||||
# The logical name of the MCP server a tool call was routed to. There is no
|
||||
# semconv key for an MCP server's *name* (the convention uses ``server.address``
|
||||
|
|
|
|||
|
|
@ -198,10 +198,58 @@ def guardrail_span_name(data: "GuardrailSpanData") -> str:
|
|||
return f"execute_guardrail {data.guardrail_name}".strip()
|
||||
|
||||
|
||||
_SERVICE_VERB_BY_CALL_TYPE: Final[dict[str, str]] = {
|
||||
"get_cache": "get",
|
||||
"async_get_cache": "get",
|
||||
"batch_get_cache": "mget",
|
||||
"async_batch_get_cache": "mget",
|
||||
"set_cache": "set",
|
||||
"async_set_cache": "set",
|
||||
"async_set_cache_pipeline": "set",
|
||||
"async_set_cache_pipeline_with_ttls": "set",
|
||||
"async_set_cache_sadd": "sadd",
|
||||
"increment_cache": "incr",
|
||||
"async_increment": "incr",
|
||||
"async_increment_pipeline": "incr",
|
||||
"delete_cache": "delete",
|
||||
"async_delete_cache": "delete",
|
||||
"async_rpush": "rpush",
|
||||
"async_lpop": "lpop",
|
||||
"async_scan_iter": "scan",
|
||||
"async_lpop_pipeline": "lpop",
|
||||
"async_rpush_pipeline": "rpush",
|
||||
"async_rpush_and_trim": "rpush",
|
||||
"increment_cache_ttl": "ttl",
|
||||
"increment_cache_expire": "expire",
|
||||
"async_ping": "ping",
|
||||
"sync_ping": "ping",
|
||||
"redis_async_ping": "ping",
|
||||
"redis_sync_ping": "ping",
|
||||
"request_redis_batch": "pipeline",
|
||||
"post_call_redis_batch": "pipeline",
|
||||
}
|
||||
|
||||
|
||||
def service_operation(data: "ServiceSpanData") -> str | None:
|
||||
"""``"redis.get"`` when the call type is a known datastore verb, else ``None``
|
||||
(Postgres helpers stay function-named until they get ``db.select {table}`` names)."""
|
||||
if not data.call_type:
|
||||
return None
|
||||
verb: Final = _SERVICE_VERB_BY_CALL_TYPE.get(data.call_type)
|
||||
if verb is None:
|
||||
return None
|
||||
return f"{data.service_name}.{verb}"
|
||||
|
||||
|
||||
def service_span_name(data: "ServiceSpanData") -> str:
|
||||
"""``"{service} {call_type}"`` e.g. ``"redis set"`` — service name alone when
|
||||
no call type is known, so identically-named calls stay distinguishable."""
|
||||
return f"{data.service_name} {data.call_type or ''}".strip()
|
||||
"""``"{service}.{verb} {target}"`` (``"redis.get llm_response"``) for a known datastore
|
||||
verb, ``"{service}.{verb}"`` (``"redis.pipeline"``) when the producer declared no
|
||||
target, else ``"{service} {call_type}"`` (``"postgres get_data"``) — service name alone
|
||||
when no call type is known, so identically-named calls stay distinguishable."""
|
||||
operation: Final = service_operation(data)
|
||||
if operation is None:
|
||||
return f"{data.service_name} {data.call_type or ''}".strip()
|
||||
return f"{operation} {data.target}" if data.target else operation
|
||||
|
||||
|
||||
def root_roles() -> list[SpanRole]:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
"""Trace-context + Baggage helpers."""
|
||||
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar, Token
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
|
|
@ -246,16 +247,53 @@ def resolve_service_span_context(
|
|||
return set_span_in_context(INVALID_SPAN, ctx), (Link(parent.get_span_context()),)
|
||||
|
||||
|
||||
_post_response_root: Final["ContextVar[SpanContext | None]"] = ContextVar(
|
||||
"litellm_otel_post_response_root", default=None
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def post_response_root(span: Span) -> Generator[None]:
|
||||
"""Nest the post-response service calls inside this block under ``span``."""
|
||||
token: Final = _post_response_root.set(span.get_span_context())
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_post_response_root.reset(token)
|
||||
|
||||
|
||||
def _is_post_response(parent: Span, end_time_ns: int | None) -> bool:
|
||||
if not isinstance(parent, ReadableSpan):
|
||||
return False
|
||||
if in_post_response_phase():
|
||||
return True
|
||||
return parent.get_span_context() != _post_response_root.get()
|
||||
if parent.end_time is None:
|
||||
return False
|
||||
return end_time_ns is None or end_time_ns > parent.end_time
|
||||
|
||||
|
||||
_active_phase_span: Final["ContextVar[Span | None]"] = ContextVar("litellm_otel_active_phase_span", default=None)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def active_phase(span: Span) -> Generator[None]:
|
||||
"""Make ``span`` the phase that request-level spans opened inside the block nest under.
|
||||
|
||||
A ContextVar rather than the ambient span so a close callback whose task was
|
||||
spawned inside the phase still parents to it, while one spawned after the
|
||||
phase exited sees no phase at all.
|
||||
"""
|
||||
token: Final = _active_phase_span.set(span)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_active_phase_span.reset(token)
|
||||
|
||||
|
||||
def active_phase_span() -> Span | None:
|
||||
return _active_phase_span.get()
|
||||
|
||||
|
||||
def resolve_request_span_context() -> Context:
|
||||
"""The parent context for a request-level span (the LLM call, a guardrail).
|
||||
|
||||
|
|
@ -267,7 +305,7 @@ def resolve_request_span_context() -> Context:
|
|||
|
||||
Unlike :func:`resolve_parent_context` (used by DB/service spans, which DO want
|
||||
to nest under the active phase span, e.g. an auth DB lookup under ``auth``),
|
||||
this never returns the active span when an anchor exists.
|
||||
this never returns the momentarily active span when an anchor exists.
|
||||
"""
|
||||
root: Final = request_root_span()
|
||||
if root is not None:
|
||||
|
|
@ -275,6 +313,20 @@ def resolve_request_span_context() -> Context:
|
|||
return get_current()
|
||||
|
||||
|
||||
def resolve_internal_call_span_context() -> Context:
|
||||
"""The parent context for an LLM call litellm itself makes while working a request.
|
||||
|
||||
The auto-router classifier runs inside ``route {model_group}``; that phase, opened
|
||||
with :func:`active_phase`, owns the sub-call so it reads as part of routing rather
|
||||
than as a second provider attempt beside the caller's own ``chat``. With no phase
|
||||
open the sub-call anchors like any request-level span.
|
||||
"""
|
||||
phase: Final = active_phase_span()
|
||||
if phase is not None:
|
||||
return context_from_span(phase)
|
||||
return resolve_request_span_context()
|
||||
|
||||
|
||||
def resolve_mcp_span_context(
|
||||
carrier: "Mapping[str, str] | None" = None,
|
||||
) -> "tuple[Context, tuple[Link, ...]]":
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from pydantic import BaseModel
|
|||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import PROXY_LLM_PROVIDER_FALLBACK, PROXY_REJECTED_BEFORE_ROUTING_KEY
|
||||
from litellm.exceptions import (
|
||||
|
|
@ -45,6 +46,7 @@ from litellm.proxy._types import (
|
|||
LiteLLM_UserTable,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
|
||||
from litellm.repositories.base_repository import BaseRepository
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
|
|
@ -4182,6 +4184,7 @@ class PrometheusLogger(CustomLogger):
|
|||
self._get_remaining_hours_for_budget_reset(budget_reset_at=budget_reset_at)
|
||||
)
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _set_customer_budget_metrics_after_api_request(
|
||||
self,
|
||||
end_user_id: str | None,
|
||||
|
|
|
|||
|
|
@ -5198,13 +5198,18 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
|
|||
Returns ``None`` when V2 is off OR when there's no preset registered for
|
||||
``callback_name`` — callers should then fall through to the legacy path.
|
||||
|
||||
A preset that needs operator credentials it cannot find is allowed to build
|
||||
only when this request has a key/team destination for that backend and another
|
||||
V2 logger is already registered to carry the fan-out. The resulting logger keeps
|
||||
only its credential-gated exporter, while the registered logger owns operator
|
||||
delivery. Without that carrier, a preset that raises or that ends up with nothing
|
||||
but its gated exporter and the default console placeholder returns ``None``, so the
|
||||
caller falls through to the legacy path exactly as before V2 landed.
|
||||
A logger built while another V2 logger is already registered keeps only the
|
||||
exporters its own preset contributed, whether or not the operator holds
|
||||
credentials for that backend and whether or not a destination is anchored: the
|
||||
registered logger owns operator delivery, so a copy of the operator's base OTLP
|
||||
exporters here would emit every LLM call a second time into the operator's sink.
|
||||
A preset that contributes no exporter of its own (a mapper over the operator's
|
||||
collector) keeps the base exporters, since it has nothing else to deliver through.
|
||||
A preset that needs operator credentials it cannot find is allowed to build only
|
||||
when it serves a key/team destination in that situation. Otherwise a preset that
|
||||
raises or that ends up with nothing but its gated exporter and the default
|
||||
console placeholder returns ``None``, so the caller falls through to the legacy
|
||||
path exactly as before V2 landed.
|
||||
"""
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
||||
|
|
@ -5236,7 +5241,7 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
|
|||
gated: Final = _is_credential_gated(built)
|
||||
if gated and not carried and not _has_operator_exporter(built):
|
||||
return None
|
||||
config: Final = _only_the_gated_exporter(built) if gated and carried else built
|
||||
config: Final = _only_the_presets_own_exporters(built, callback_name) if has_v2_logger else built
|
||||
if _exports_nowhere(config):
|
||||
verbose_logger.warning(
|
||||
"OTel V2: no operator credentials for '%s'; only key/team destinations will receive its traces",
|
||||
|
|
@ -5264,8 +5269,10 @@ def _has_operator_exporter(config: "OpenTelemetryV2Config") -> bool:
|
|||
return any(not _is_gated(spec) and not is_unconfigured_placeholder(spec) for spec in config.exporters)
|
||||
|
||||
|
||||
def _only_the_gated_exporter(config: "OpenTelemetryV2Config") -> "OpenTelemetryV2Config":
|
||||
return config.model_copy(update={"exporters": [spec for spec in config.exporters if _is_gated(spec)]})
|
||||
def _only_the_presets_own_exporters(config: "OpenTelemetryV2Config", callback_name: str) -> "OpenTelemetryV2Config":
|
||||
"""A preset with no exporter of its own (Langtrace: a mapper over the operator's collector) keeps the base."""
|
||||
own: Final = [spec for spec in config.exporters if spec.owner == callback_name]
|
||||
return config.model_copy(update={"exporters": own}) if own else config
|
||||
|
||||
|
||||
def _is_gated(spec: "ExporterSpec") -> bool:
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from collections.abc import Mapping
|
|||
from typing import TYPE_CHECKING, Any, Final, Generic, Protocol, TypeVar, cast, runtime_checkable
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm.llms.base_llm.managed_resources.isolation import (
|
||||
build_list_page,
|
||||
build_owner_filter,
|
||||
|
|
@ -18,6 +19,8 @@ from litellm.llms.base_llm.managed_resources.isolation import (
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
MANAGED_RESOURCES_TARGET: Final = "managed_resources"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
|
|
@ -158,6 +161,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
# COMMON STORAGE OPERATIONS
|
||||
# ============================================================================
|
||||
|
||||
@with_service_target(MANAGED_RESOURCES_TARGET)
|
||||
async def store_unified_resource_id(
|
||||
self,
|
||||
unified_resource_id: str,
|
||||
|
|
@ -240,6 +244,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
"LiteLLM Managed %s with id=%s stored in db: %s", self.resource_type, unified_resource_id, result
|
||||
)
|
||||
|
||||
@with_service_target(MANAGED_RESOURCES_TARGET)
|
||||
async def get_unified_resource_id(
|
||||
self,
|
||||
unified_resource_id: str,
|
||||
|
|
@ -276,6 +281,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]):
|
|||
|
||||
return None
|
||||
|
||||
@with_service_target(MANAGED_RESOURCES_TARGET)
|
||||
async def delete_unified_resource_id(
|
||||
self,
|
||||
unified_resource_id: str,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from starlette.types import Scope
|
|||
from typing_extensions import assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_ALL_TOOLS_WILDCARD
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -60,6 +61,7 @@ from litellm.proxy.auth.user_api_key_auth import (
|
|||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
USER_NO_MCP_PERMISSION_SENTINEL,
|
||||
get_management_object_ttl,
|
||||
user_object_permission_id_cache_key,
|
||||
|
|
@ -3091,6 +3093,7 @@ class MCPRequestHandler:
|
|||
return object_permission
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _user_object_permission_id(
|
||||
user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False
|
||||
) -> str | None:
|
||||
|
|
@ -3395,6 +3398,7 @@ class MCPRequestHandler:
|
|||
_AGENT_NO_PERMISSION_SENTINEL = "__agent_no_mcp_permission__"
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _agent_object_permission_id(agent_id: str, prisma_client: "PrismaClient") -> str | None:
|
||||
"""The permission row this agent's row links to, or ``None`` when it links none.
|
||||
|
||||
|
|
|
|||
|
|
@ -53,6 +53,7 @@ from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Resp
|
|||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -93,6 +94,8 @@ from litellm.proxy.common_utils.html_forms.native_client_consent import (
|
|||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
_DCR_CLAIMS_TARGET: Final = "mcp_dcr_claims"
|
||||
|
||||
GATEWAY_DCR_CLIENT_ID_PREFIX: Final = "llm_dcrc_"
|
||||
"""Marker prefix on every gateway-issued DCR client_id so the root authorize/token
|
||||
endpoints can route an aggregate-flow request without decrypting, and existing per-server
|
||||
|
|
@ -1017,6 +1020,7 @@ class _SingleUseGuard:
|
|||
def __init__(self, cache: DualCache) -> None:
|
||||
self._cache = cache
|
||||
|
||||
@with_service_target(_DCR_CLAIMS_TARGET)
|
||||
async def claim(self, key: str, ttl_seconds: int) -> ClaimOutcome:
|
||||
"""Atomically claim ``key``. ``"first"`` iff this caller is the first (increment to 1),
|
||||
``"replayed"`` on a replay (>1), and ``"unavailable"`` when the claim could not be recorded in
|
||||
|
|
@ -1045,6 +1049,7 @@ class _SingleUseGuard:
|
|||
count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds, local_only=True)
|
||||
return "first" if count == 1 else "replayed"
|
||||
|
||||
@with_service_target(_DCR_CLAIMS_TARGET)
|
||||
async def peek(self, key: str) -> Literal["unclaimed", "claimed", "unavailable"]:
|
||||
"""Read-only view of a single-use marker, resolved against the same shared authority as
|
||||
:meth:`claim` so introspection observes exactly the record redemption and revocation wrote.
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ from pydantic import AnyUrl, BaseModel, TypeAdapter
|
|||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -150,6 +151,7 @@ from litellm.proxy._experimental.mcp_server.upstream import (
|
|||
to_server_spec_fail_closed,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_SERVERS_TARGET,
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
MCPMissingUserEnvVarsError,
|
||||
add_server_prefix_to_name,
|
||||
|
|
@ -3299,6 +3301,7 @@ class MCPServerManager:
|
|||
def get_byom_submitted_servers_cache_key(user_id: str) -> str:
|
||||
return f"byom_submitted_servers:{user_id}"
|
||||
|
||||
@with_service_target(MCP_SERVERS_TARGET)
|
||||
async def invalidate_byom_submitted_servers_cache(self, user_id: str | None) -> None:
|
||||
if not user_id:
|
||||
return
|
||||
|
|
@ -3309,6 +3312,7 @@ class MCPServerManager:
|
|||
except Exception as e: # noqa: BLE001
|
||||
verbose_logger.warning("Failed to invalidate BYOM submitted MCP server cache: %s", e)
|
||||
|
||||
@with_service_target(MCP_SERVERS_TARGET)
|
||||
async def _get_active_submitted_mcp_server_ids_for_user(
|
||||
self, user_api_key_auth: UserAPIKeyAuth | None
|
||||
) -> list[str]:
|
||||
|
|
@ -3542,6 +3546,7 @@ class MCPServerManager:
|
|||
if not explicit_grants_only and (scope is None or server_id == scope)
|
||||
]
|
||||
|
||||
@with_service_target(MCP_SERVERS_TARGET)
|
||||
async def resolve_toolset_tool_permissions(
|
||||
self,
|
||||
toolset_ids: list[str],
|
||||
|
|
@ -3636,6 +3641,7 @@ class MCPServerManager:
|
|||
except Exception as e:
|
||||
verbose_logger.warning("invalidate_toolset_cache: failed to evict in-memory entries: %s", e)
|
||||
|
||||
@with_service_target(MCP_SERVERS_TARGET)
|
||||
async def get_toolset_by_name_cached(
|
||||
self,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Final
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -30,6 +31,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import OAuthTokenCacheCodec
|
||||
from litellm.proxy._experimental.mcp_server.utils import MCP_OAUTH_TOKENS_TARGET
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
|
|
@ -245,6 +247,7 @@ class MCPPerUserTokenCache:
|
|||
token: Final = await self.get_token(user_id, server_id)
|
||||
return token.access_token if token is not None else None
|
||||
|
||||
@with_service_target(MCP_OAUTH_TOKENS_TARGET)
|
||||
async def get_token(self, user_id: str, server_id: str) -> OAuthToken | None:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415
|
||||
|
|
@ -263,6 +266,7 @@ class MCPPerUserTokenCache:
|
|||
)
|
||||
return None
|
||||
|
||||
@with_service_target(MCP_OAUTH_TOKENS_TARGET)
|
||||
async def set(
|
||||
self,
|
||||
user_id: str,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from __future__ import annotations
|
|||
from dataclasses import KW_ONLY, dataclass
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
OAuthToken,
|
||||
|
|
@ -19,6 +20,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
|
|||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
|
||||
OAuthTokenCacheCodec,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import MCP_OAUTH_TOKENS_TARGET
|
||||
|
||||
|
||||
class AsyncCache(Protocol):
|
||||
|
|
@ -47,6 +49,7 @@ class DualCacheTokenCacheBackend:
|
|||
def _key(self, user_id: str, server_id: str) -> str:
|
||||
return f"{self.key_prefix}{user_id}:{server_id}"
|
||||
|
||||
@with_service_target(MCP_OAUTH_TOKENS_TARGET)
|
||||
async def get(self, user_id: str, server_id: str) -> OAuthToken | None:
|
||||
try:
|
||||
blob: Final = await self.cache.async_get_cache(self._key(user_id, server_id))
|
||||
|
|
@ -55,6 +58,7 @@ class DualCacheTokenCacheBackend:
|
|||
verbose_logger.debug("MCP per-user token cache get failed (miss): %s", exc)
|
||||
return None
|
||||
|
||||
@with_service_target(MCP_OAUTH_TOKENS_TARGET)
|
||||
async def set(self, user_id: str, server_id: str, token: OAuthToken, ttl_seconds: float) -> None:
|
||||
if ttl_seconds <= 0:
|
||||
return
|
||||
|
|
@ -67,6 +71,7 @@ class DualCacheTokenCacheBackend:
|
|||
except Exception as exc: # noqa: BLE001
|
||||
verbose_logger.debug("MCP per-user token cache set failed (ignored): %s", exc)
|
||||
|
||||
@with_service_target(MCP_OAUTH_TOKENS_TARGET)
|
||||
async def delete(self, user_id: str, server_id: str) -> None:
|
||||
try:
|
||||
await self.cache.async_delete_cache(self._key(user_id, server_id))
|
||||
|
|
|
|||
|
|
@ -19,6 +19,9 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
|||
if typing.TYPE_CHECKING:
|
||||
from fastapi import Request
|
||||
|
||||
MCP_SERVERS_TARGET: Final = "mcp_servers"
|
||||
MCP_OAUTH_TOKENS_TARGET: Final = "mcp_oauth_tokens"
|
||||
|
||||
|
||||
class _McpServerLike(Protocol):
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from collections.abc import Mapping
|
|||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm.proxy.agent_endpoints.managed_identity import classify_agent_subject
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, get_management_object_ttl
|
||||
from litellm.repositories.table_repositories import (
|
||||
|
|
@ -20,6 +21,8 @@ from litellm.types.proxy.agent_identity import (
|
|||
VerifiedHumanSubject,
|
||||
)
|
||||
|
||||
_AGENT_IDENTITIES_TARGET: Final = "agent_identities"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma.models import LiteLLM_VerifiedSubject
|
||||
from prisma.types import (
|
||||
|
|
@ -90,6 +93,7 @@ class AgentIdentityStore:
|
|||
return AgentIdentityFailure(message="This agent identity binding has been retired")
|
||||
return None
|
||||
|
||||
@with_service_target(_AGENT_IDENTITIES_TARGET)
|
||||
async def _bound_agent_id(self, tenant_id: str, client_id: str) -> str | AgentIdentityFailure | None:
|
||||
cache_key: Final = f"agent_identity:{json.dumps((tenant_id, client_id))}"
|
||||
cached: Final[object] = await self.cache.async_get_cache(key=cache_key) if self.cache is not None else None
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from fastapi import APIRouter, Depends, Request, Response
|
|||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel, Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -37,6 +38,7 @@ from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
|
|||
from litellm.proxy.anthropic_endpoints.endpoints import anthropic_response, count_tokens
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_set_request_parsed_body
|
||||
from litellm.proxy.management_endpoints.sso_helper_utils import CLI_SSO_SESSIONS_TARGET
|
||||
from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail
|
||||
|
||||
GATEWAY_PREFIX: Final = "/claude_code_gateway"
|
||||
|
|
@ -271,6 +273,7 @@ def _mint_access_token(login: _GatewayLogin) -> str:
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(CLI_SSO_SESSIONS_TARGET)
|
||||
async def _claim_device_code(login_id: str, cache: DualCache) -> bool:
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_get_cli_sso_flow_cache_key, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
|
|
@ -284,6 +287,7 @@ async def _claim_device_code(login_id: str, cache: DualCache) -> bool:
|
|||
return claims == 1
|
||||
|
||||
|
||||
@with_service_target(CLI_SSO_SESSIONS_TARGET)
|
||||
async def _handle_device_code_grant(device_code: str | None) -> JSONResponse:
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
|
|||
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache, LimitedSizeOrderedDict
|
||||
from litellm.constants import (
|
||||
|
|
@ -94,6 +95,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
from litellm.proxy.common_utils.model_listing_utils import alias_map
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
|
||||
MODEL_ACCESS_GROUP_REGISTRY_OVERFLOW_SENTINEL,
|
||||
NO_TEAM_MEMBERSHIP_SENTINEL,
|
||||
|
|
@ -122,6 +124,7 @@ from litellm.proxy.guardrails.tool_name_extraction import (
|
|||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
|
||||
from litellm.proxy.spend_tracking.carried_budget_state import carry_organization_budget_state
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import SPEND_COUNTERS_TARGET
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
|
||||
|
|
@ -1443,6 +1446,7 @@ def get_key_end_user_budget_id(key_metadata: Mapping[str, object] | None) -> str
|
|||
return budget_id if isinstance(budget_id, str) and budget_id != "" else None
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_default_end_user_budget(
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -1509,6 +1513,7 @@ async def get_default_end_user_budget(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_team_member_default_budget(
|
||||
budget_id: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -1710,6 +1715,7 @@ _END_USER_REGISTRY_LOAD_LOCK: Final = asyncio.Lock()
|
|||
_MODEL_ACCESS_GROUP_REGISTRY_LOAD_LOCK: Final = asyncio.Lock()
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _cached_registry(
|
||||
cache_key: str,
|
||||
overflow_sentinel: str,
|
||||
|
|
@ -1725,6 +1731,7 @@ async def _cached_registry(
|
|||
return _REGISTRY_NOT_CACHED
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _cache_registry_answer(
|
||||
cache_key: str,
|
||||
value: tuple[str, ...] | str,
|
||||
|
|
@ -1873,6 +1880,7 @@ async def _end_user_is_known_unrestricted(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_end_user_object(
|
||||
end_user_id: str | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -1978,6 +1986,7 @@ _END_USER_VALIDATION_NEGATIVE_TTL: Final = 60
|
|||
_END_USER_VALIDATION_POSITIVE_TTL: Final = 300
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def resolve_and_validate_end_user_id(
|
||||
raw_end_user_id: str | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -2127,6 +2136,7 @@ async def _load_model_access_group_registry(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _fetch_uncached_model_access_group_budgets(
|
||||
uncached_groups: Sequence[str],
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -2173,6 +2183,7 @@ def _model_access_group_budget(row: _PrismaModelAccessGroupBudgetRow) -> ModelAc
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_model_access_group_budgets_batch(
|
||||
access_group_names: Sequence[str],
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -2204,6 +2215,7 @@ async def get_model_access_group_budgets_batch(
|
|||
return {group: budget for group, budget in (*probed, *fetched) if budget is not None}
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _fetch_uncached_tags(
|
||||
uncached_tags: Sequence[str],
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -2244,6 +2256,7 @@ async def _fetch_uncached_tags(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_tag_objects_batch(
|
||||
tag_names: Sequence[str],
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -2337,6 +2350,7 @@ def _membership_from_cached_payload(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _fetch_team_membership_from_db(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
|
|
@ -2367,6 +2381,7 @@ async def _fetch_team_membership_from_db(
|
|||
return membership
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _load_team_membership_on_cache_miss(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
|
|
@ -2391,6 +2406,7 @@ async def _load_team_membership_on_cache_miss(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_team_membership(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
|
|
@ -2584,6 +2600,7 @@ async def _get_fuzzy_user_object(
|
|||
return response
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _backfill_null_user_email(
|
||||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -2613,6 +2630,7 @@ async def _backfill_null_user_email(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_user_object(
|
||||
user_id: str | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -2768,6 +2786,7 @@ def _user_read_failure(user_id: str, error: Exception) -> Exception:
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _cache_management_object(
|
||||
key: str,
|
||||
value: BaseModel | Mapping[str, object],
|
||||
|
|
@ -2789,6 +2808,7 @@ async def _cache_management_object(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _cache_team_object(
|
||||
team_id: str,
|
||||
team_table: LiteLLM_TeamTableCachedObj,
|
||||
|
|
@ -2846,6 +2866,7 @@ async def _cache_team_object(
|
|||
await _invalidate_usage_cache_entry(usage_cache, alias_key, redis_shared=redis_shared, stale="team alias")
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _invalidate_usage_cache_entry(
|
||||
usage_cache: DualCache | None,
|
||||
key: str,
|
||||
|
|
@ -2869,6 +2890,7 @@ async def _invalidate_usage_cache_entry(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def invalidate_team_member_spend_state(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
|
|
@ -2985,6 +3007,7 @@ async def invalidate_team_member_spend_state(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def delete_cache_team_object(
|
||||
team_id: str,
|
||||
team_alias: str | None,
|
||||
|
|
@ -3044,6 +3067,7 @@ async def _cache_key_object(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _delete_cache_key_object(
|
||||
hashed_token: str,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -3241,6 +3265,7 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
return _response
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _get_team_object_from_cache(
|
||||
key: str,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -3316,6 +3341,7 @@ async def get_team_object(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _cache_access_object(
|
||||
access_group_id: str,
|
||||
access_group_table: LiteLLM_AccessGroupTable,
|
||||
|
|
@ -3331,6 +3357,7 @@ async def _cache_access_object(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _delete_cache_access_object(
|
||||
access_group_id: str,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -3346,6 +3373,7 @@ async def _delete_cache_access_object(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_access_object(
|
||||
access_group_id: str,
|
||||
prisma_client: DatabaseClient | None,
|
||||
|
|
@ -3417,6 +3445,7 @@ async def get_access_object(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_team_object_by_alias(
|
||||
team_alias: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -3527,6 +3556,7 @@ async def get_team_object_by_alias(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_org_object_by_alias(
|
||||
org_alias: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -3882,6 +3912,7 @@ async def get_jwt_key_mapping_object(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_key_object(
|
||||
hashed_token: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -3979,6 +4010,7 @@ def _copy_user_api_key_auth_for_cache(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_object_permission(
|
||||
object_permission_id: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -4035,6 +4067,7 @@ async def get_object_permission(
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_managed_vector_store_rows_by_uuids(
|
||||
uuids: list[str],
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -4104,6 +4137,7 @@ class OrganizationNotFoundError(Exception):
|
|||
|
||||
|
||||
@log_db_metrics
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_org_object(
|
||||
org_id: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
@ -4180,6 +4214,7 @@ def _last_known_org_cache_key(org_id: str) -> str:
|
|||
return f"org_id:{org_id}:with_budget:last_known"
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _keep_last_known_org(
|
||||
org: LiteLLM_OrganizationTable, org_id: str, user_api_key_cache: UserApiKeyCache
|
||||
) -> None:
|
||||
|
|
@ -4197,6 +4232,7 @@ async def _keep_last_known_org(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_org_object_for_request(
|
||||
org_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -6235,6 +6271,7 @@ async def _project_soft_budget_check(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_project_object(
|
||||
project_id: str,
|
||||
prisma_client: PrismaClient | None,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import Final, Literal, Protocol, TypeAlias
|
|||
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_batch import active_request_redis_batch
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
|
@ -23,6 +24,7 @@ from litellm.models.user import LiteLLM_UserTable
|
|||
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
UserApiKeyCache,
|
||||
get_management_object_ttl,
|
||||
team_membership_auth_cache_key,
|
||||
|
|
@ -222,13 +224,14 @@ def _set_in_memory(memory: _InMemoryCache, cache_key: str, value: object, ttl: f
|
|||
async def _read_redis_rows(keys: list[str], redis_cache: RedisCache) -> Mapping[str, object]:
|
||||
"""On the request pipeline when one is open; a failed pipeline reads as a miss, like ``async_batch_get_cache``."""
|
||||
batch: Final = active_request_redis_batch(redis_cache)
|
||||
if batch is None:
|
||||
return await redis_cache.async_batch_get_cache(key_list=keys) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API
|
||||
try:
|
||||
return await batch.mget(keys)
|
||||
except Exception as e: # noqa: BLE001 # the DB fill below takes over, as it does after a failed MGET today
|
||||
verbose_proxy_logger.debug("auth prefetch Redis read failed, filling from the database: %s", e)
|
||||
return MappingProxyType({})
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
if batch is None:
|
||||
return await redis_cache.async_batch_get_cache(key_list=keys) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API
|
||||
try:
|
||||
return await batch.mget(keys)
|
||||
except Exception as e: # noqa: BLE001 # the DB fill below takes over, as it does after a failed MGET today
|
||||
verbose_proxy_logger.debug("auth prefetch Redis read failed, filling from the database: %s", e)
|
||||
return MappingProxyType({})
|
||||
|
||||
|
||||
async def _fill_from_redis(entries: Sequence[_CacheEntry], redis_cache: RedisCache, memory: _InMemoryCache) -> None:
|
||||
|
|
@ -283,11 +286,12 @@ async def _write_back(entries: Sequence[tuple[_CacheEntry, BaseModel]], cache: U
|
|||
if cache.redis_cache is None:
|
||||
return
|
||||
batch: Final = active_request_redis_batch(cache.redis_cache)
|
||||
if batch is None:
|
||||
await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads)
|
||||
return
|
||||
for cache_key, payload, ttl in payloads: # rides the request's next round trip; the scope drains leftovers
|
||||
batch.set(cache_key, payload, ttl)
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
if batch is None:
|
||||
await cache.redis_cache.async_set_cache_pipeline_with_ttls(payloads)
|
||||
return
|
||||
for cache_key, payload, ttl in payloads: # rides the request's next round trip; the scope drains leftovers
|
||||
batch.set(cache_key, payload, ttl)
|
||||
|
||||
|
||||
async def _fill_from_db(
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from fastapi import HTTPException, status
|
|||
from jwt.api_jwk import PyJWK
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
|
||||
from litellm.llms.custom_httpx.httpx_handler import HTTPHandler
|
||||
|
|
@ -65,6 +66,7 @@ from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canon
|
|||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
UserApiKeyCache,
|
||||
get_management_object_ttl,
|
||||
)
|
||||
|
|
@ -783,15 +785,18 @@ class JWTHandler:
|
|||
except httpx.TransportError as e:
|
||||
raise JWKSUnreachableError(f"{type(e).__name__} fetching {url} after {JWKS_FETCH_ATTEMPTS} attempts") from e
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _get_cached_value(self, cache_key: str) -> _CachedValueT | None:
|
||||
cached: Final = await self.user_api_key_cache.async_get_cache(cache_key)
|
||||
return cast("_CachedValueT | None", cached) # cast-ok: cache reads are untyped
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _get_cached_timestamp(self, cache_key: str) -> float | None:
|
||||
cached: Final = await self.user_api_key_cache.async_get_cache(cache_key)
|
||||
# A JSON round-trip through Redis hands a whole-number epoch back as an int.
|
||||
return float(cached) if isinstance(cached, (int, float)) else None
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _put_cached_value(self, cache_key: str, value: JWKKeyValue | str | float, ttl: float) -> None:
|
||||
await self.user_api_key_cache.async_set_cache(key=cache_key, value=value, ttl=ttl)
|
||||
|
||||
|
|
@ -1006,6 +1011,7 @@ class JWTHandler:
|
|||
else:
|
||||
return False
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def get_oidc_userinfo(self, token: str) -> dict:
|
||||
"""
|
||||
Fetch user information from OIDC UserInfo endpoint.
|
||||
|
|
@ -2057,6 +2063,7 @@ class JWTAuthManager:
|
|||
return
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def sync_user_role_and_teams(
|
||||
jwt_handler: JWTHandler,
|
||||
jwt_valid_token: dict,
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from fastapi import Request, status
|
|||
from pydantic import TypeAdapter, ValidationError
|
||||
from redis.exceptions import RedisError
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError
|
||||
|
|
@ -38,6 +39,8 @@ from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
|||
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_client_ip
|
||||
from litellm.secret_managers.main import get_secret_bool
|
||||
|
||||
_LOGIN_THROTTLE_TARGET: Final = "login_throttle"
|
||||
|
||||
DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 10
|
||||
DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 60
|
||||
DEFAULT_FAILED_LOGIN_BLOCK_SECONDS: Final = 300
|
||||
|
|
@ -346,6 +349,7 @@ class LoginThrottle:
|
|||
return Block(scope="user", retry_after=user_ttl)
|
||||
return None
|
||||
|
||||
@with_service_target("login_throttle")
|
||||
async def _shared_block_ttls(self, keys: _Keys) -> _BlockTtls:
|
||||
if self.redis_cache is None:
|
||||
return LOGIN_THROTTLE_NOT_BLOCKED
|
||||
|
|
@ -366,6 +370,7 @@ class LoginThrottle:
|
|||
return 0
|
||||
return max(math.ceil(expires_at - time.time()), 0)
|
||||
|
||||
@with_service_target("login_throttle")
|
||||
async def record_failure(self, username: str) -> _BlockTtls:
|
||||
keys: Final = self._keys(username)
|
||||
source_limit: Final = self.source_limit or 0
|
||||
|
|
@ -393,6 +398,7 @@ class LoginThrottle:
|
|||
self.blocks.set_cache(block_key, time.time() + self.block_seconds, ttl=self.block_seconds)
|
||||
return self.block_seconds
|
||||
|
||||
@with_service_target(_LOGIN_THROTTLE_TARGET)
|
||||
async def clear_pair(self, username: str) -> None:
|
||||
pair_counter: Final = self._keys(username).pair_counter
|
||||
if self.redis_cache is not None:
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Final
|
|||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
|
|
@ -32,6 +33,7 @@ from litellm.proxy.auth.resolvers.models import (
|
|||
UserIdentity,
|
||||
)
|
||||
from litellm.proxy.auth.roles import TeamRole, map_role, team_role
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -99,6 +101,7 @@ class IdentityStore:
|
|||
raise PrincipalMissingSourceKeyError()
|
||||
return principal.source_key
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _resolve_key(self, hashed_token: str) -> UserAPIKeyAuth:
|
||||
if self._prisma is None:
|
||||
raise NoDatabaseConnectionError()
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from fastapi.security.api_key import APIKeyHeader
|
|||
from starlette.exceptions import WebSocketException
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
|
@ -73,7 +74,12 @@ from litellm.proxy.auth.auth_checks import (
|
|||
)
|
||||
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
|
||||
from litellm.proxy.auth.auth_method import AuthMethod
|
||||
from litellm.proxy.auth.auth_object_prefetch import AuthObjectRefs, prefetch_auth_objects, prefetch_identity_keys
|
||||
from litellm.proxy.auth.auth_object_prefetch import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
AuthObjectRefs,
|
||||
prefetch_auth_objects,
|
||||
prefetch_identity_keys,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
abbreviate_api_key,
|
||||
get_end_user_id_from_request_body,
|
||||
|
|
@ -3497,8 +3503,13 @@ async def user_api_key_auth(
|
|||
|
||||
# Run the whole auth phase inside a live ``auth`` span so the DB lookups it
|
||||
# triggers (key/user/team object reads) nest under it instead of flattening
|
||||
# onto the server span. No-op when OTel V2 isn't active.
|
||||
with phase_span(f"auth {route}"), spend_counter_batch_scope(_spend_counter_redis_cache()):
|
||||
# onto the server span, and name every cache read in it an auth-object read.
|
||||
# No-op when OTel V2 isn't active.
|
||||
with (
|
||||
phase_span(f"auth {route}"),
|
||||
service_target(AUTH_OBJECTS_TARGET),
|
||||
spend_counter_batch_scope(_spend_counter_redis_cache()),
|
||||
):
|
||||
try:
|
||||
user_api_key_auth_obj: Final = await _user_api_key_auth_builder(
|
||||
request=request,
|
||||
|
|
|
|||
|
|
@ -4,12 +4,14 @@ from collections.abc import Sequence
|
|||
from dataclasses import asdict, dataclass
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||
_ConfigSyncPubSub,
|
||||
_pubsub_capable_client,
|
||||
coordination_redis_cache,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
|
|
@ -123,6 +125,7 @@ async def publish_auth_cache_invalidation(
|
|||
await asyncio.sleep(0)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "UserApiKeyCache") -> None:
|
||||
"""
|
||||
Drop cached management objects here and on every other worker.
|
||||
|
|
@ -204,6 +207,7 @@ class AuthCacheInvalidationSubscriber:
|
|||
continue
|
||||
self._apply_message(message)
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
def _apply_message(self, message: object) -> None:
|
||||
data: Final = message.get("data") if isinstance(message, dict) else None
|
||||
parsed: Final = _message_from_data(data)
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import Final, Generic, Literal, Protocol, TypeVar
|
|||
from typing_extensions import assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -37,6 +38,7 @@ from litellm.proxy.common_utils.timezone_utils import (
|
|||
get_budget_reset_settings,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
end_user_cache_key,
|
||||
model_access_group_cache_key,
|
||||
model_access_group_spend_counter_key,
|
||||
|
|
@ -45,8 +47,9 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
tag_cache_key,
|
||||
)
|
||||
from litellm.proxy.db.budget_window_spend_writer import roll_window_spend_row
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import POD_LOCK_TARGET, PodLockManager
|
||||
from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import SPEND_COUNTERS_TARGET
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.prisma_protocols import PrismaBatch, SpendLinkedTable
|
||||
|
|
@ -473,6 +476,7 @@ class ResetBudgetJob:
|
|||
new_batch: Final[Callable[[], PrismaBatch]] = self.prisma_client.db.batch_
|
||||
return new_batch
|
||||
|
||||
@with_service_target(POD_LOCK_TARGET)
|
||||
async def _lease_is_held(self, lock_manager: PodLockManager) -> bool:
|
||||
"""True only when the lease is readable and someone holds it.
|
||||
|
||||
|
|
@ -570,6 +574,7 @@ class ResetBudgetJob:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _invalidate_spend_counter(counter_key: str) -> None:
|
||||
"""Drop a spend counter so the next read reseeds from the committed DB
|
||||
row, the only value that includes increments that raced the reset.
|
||||
|
|
@ -604,6 +609,7 @@ class ResetBudgetJob:
|
|||
await ResetBudgetJob._invalidate_user_api_key_cache_entry(GLOBAL_PROXY_SPEND_CACHE_KEY)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _invalidate_user_api_key_cache_entry(cache_key: str) -> None:
|
||||
"""Drop a stale management-cache entry so the next read fetches from DB.
|
||||
|
||||
|
|
@ -1373,6 +1379,7 @@ class ResetBudgetJob:
|
|||
return outcome
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _reset_expired_window(
|
||||
window: dict,
|
||||
counter_key: str,
|
||||
|
|
@ -1448,6 +1455,7 @@ class ResetBudgetJob:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _window_carried_spend(
|
||||
window: Mapping[str, object], counter_key: str, spend_counter_cache: DualCache
|
||||
) -> float:
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ if TYPE_CHECKING:
|
|||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
AUTH_OBJECTS_TARGET: Final = "auth_objects"
|
||||
|
||||
_HASHED_TOKEN_CACHE_KEY: Final = re.compile(r"[0-9a-f]{64}")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ from pydantic import TypeAdapter
|
|||
from typing_extensions import LiteralString, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import service_target, with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import RedisCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -50,7 +51,7 @@ from litellm.proxy._types import (
|
|||
SpendUpdateQueueItem,
|
||||
ToolDiscoveryQueueItem,
|
||||
)
|
||||
from litellm.proxy.common_utils.user_api_key_cache import project_cache_key
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET, project_cache_key
|
||||
from litellm.proxy.db.daily_spend_bulk_upsert import (
|
||||
DAILY_SPEND_TABLES,
|
||||
build_bulk_upsert,
|
||||
|
|
@ -2182,12 +2183,13 @@ class DBSpendUpdateWriter:
|
|||
if team_memberships_to_invalidate and proxy_logging_obj is not None:
|
||||
user_api_key_cache: Final = proxy_logging_obj.call_details.get("user_api_key_cache")
|
||||
if user_api_key_cache is not None:
|
||||
for user_id, team_id in team_memberships_to_invalidate:
|
||||
cache_key = f"team_membership:{user_id}:{team_id}"
|
||||
await user_api_key_cache.async_delete_cache(key=cache_key)
|
||||
verbose_proxy_logger.debug(
|
||||
"Invalidated team membership cache for user_id=%s, team_id=%s", user_id, team_id
|
||||
)
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
for user_id, team_id in team_memberships_to_invalidate:
|
||||
cache_key = f"team_membership:{user_id}:{team_id}"
|
||||
await user_api_key_cache.async_delete_cache(key=cache_key)
|
||||
verbose_proxy_logger.debug(
|
||||
"Invalidated team membership cache for user_id=%s, team_id=%s", user_id, team_id
|
||||
)
|
||||
elif on_table_committed is not None:
|
||||
on_table_committed("team_member_list_transactions")
|
||||
|
||||
|
|
@ -2306,6 +2308,7 @@ class DBSpendUpdateWriter:
|
|||
on_table_committed("agent_list_transactions")
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _invalidate_project_caches(project_ids: Sequence[str], proxy_logging_obj: ProxyLogging | None) -> None:
|
||||
if not project_ids or proxy_logging_obj is None:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import json
|
|||
import logging
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.redis_cache import RedisCache, log_redis_failure
|
||||
|
|
@ -10,6 +11,8 @@ from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS
|
|||
from litellm.proxy.db.db_transaction_queue.base_update_queue import service_logger_obj
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
POD_LOCK_TARGET: Final = "pod_lock"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
ProxyLogging = Any
|
||||
else:
|
||||
|
|
@ -40,6 +43,7 @@ end
|
|||
def get_redis_lock_key(cronjob_id: str) -> str:
|
||||
return f"cronjob_lock:{cronjob_id}"
|
||||
|
||||
@with_service_target(POD_LOCK_TARGET)
|
||||
async def acquire_lock(
|
||||
self,
|
||||
cronjob_id: str,
|
||||
|
|
@ -154,6 +158,7 @@ end
|
|||
except Exception as e:
|
||||
log_redis_failure(verbose_proxy_logger, logging.ERROR, f"Error releasing Redis lock for {cronjob_id}", e)
|
||||
|
||||
@with_service_target(POD_LOCK_TARGET)
|
||||
async def _compare_and_delete_lock(self, lock_key: str) -> int:
|
||||
"""
|
||||
Atomically delete lock key only if current pod owns it.
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias, TypeVar, cast
|
|||
|
||||
from redis.exceptions import RedisError
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import RedisCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -59,6 +60,8 @@ from litellm.types.caching import (
|
|||
)
|
||||
from litellm.types.services import ServiceTypes
|
||||
|
||||
SPEND_QUEUE_TARGET: Final = "spend_queue"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
else:
|
||||
|
|
@ -160,6 +163,7 @@ class RedisUpdateBuffer:
|
|||
return False
|
||||
return _use_redis_transaction_buffer
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def _store_transactions_in_redis(
|
||||
self,
|
||||
transactions: Mapping[str, BaseDailySpendTransaction] | None,
|
||||
|
|
@ -201,6 +205,7 @@ class RedisUpdateBuffer:
|
|||
str(e),
|
||||
)
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def store_in_memory_spend_updates_in_redis(
|
||||
self,
|
||||
spend_update_queue: SpendUpdateQueue,
|
||||
|
|
@ -483,6 +488,7 @@ class RedisUpdateBuffer:
|
|||
if window_spend_update_transactions and window_spend_update_queue is not None:
|
||||
await window_spend_update_queue.update_queue.put(window_spend_update_transactions)
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def restore_transactions_to_redis(
|
||||
self,
|
||||
db_spend_update_transactions: DBSpendUpdateTransactions | None = None,
|
||||
|
|
@ -543,6 +549,7 @@ class RedisUpdateBuffer:
|
|||
str(e),
|
||||
)
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def store_spend_logs_in_redis(
|
||||
self,
|
||||
rows: Sequence[SpendLogRow],
|
||||
|
|
@ -572,6 +579,7 @@ class RedisUpdateBuffer:
|
|||
verbose_proxy_logger.info("Spend tracking - parked %d spend log rows in Redis for a later flush", len(rows))
|
||||
return True
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def get_spend_logs_from_redis_buffer(self, limit: int) -> tuple[dict[str, object], ...]:
|
||||
"""Atomically take up to ``limit`` parked spend-log rows out of Redis."""
|
||||
if self.redis_cache is None or not self._should_commit_spend_updates_to_redis():
|
||||
|
|
@ -604,6 +612,7 @@ class RedisUpdateBuffer:
|
|||
"""
|
||||
return {key.replace(prefix, "", 1): value for key, value in data.items()}
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def get_all_update_transactions_from_redis_buffer(
|
||||
self,
|
||||
) -> DBSpendUpdateTransactions | None:
|
||||
|
|
@ -671,6 +680,7 @@ class RedisUpdateBuffer:
|
|||
|
||||
return combined_transaction
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def get_all_transactions_from_redis_buffer_pipeline(
|
||||
self,
|
||||
) -> tuple[
|
||||
|
|
@ -783,6 +793,7 @@ class RedisUpdateBuffer:
|
|||
service_type=ServiceTypes.REDIS_DAILY_TAG_SPEND_UPDATE_QUEUE,
|
||||
)
|
||||
|
||||
@with_service_target(SPEND_QUEUE_TARGET)
|
||||
async def _lpop_daily_spend_transactions(
|
||||
self,
|
||||
redis_key: str,
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from typing import TYPE_CHECKING, Final, TypeAlias
|
|||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import RedisCache
|
||||
from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
|
||||
|
|
@ -39,6 +40,8 @@ from litellm.types.proxy.gateway_requests import (
|
|||
GatewayRequestSnapshot,
|
||||
)
|
||||
|
||||
_GATEWAY_REQUEST_QUEUE_TARGET: Final = "gateway_request_queue"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
|
@ -166,6 +169,7 @@ class GatewayRequestRedisBuffer:
|
|||
self._redis_cache: Final = redis_cache
|
||||
self._pod_lock_manager: Final = pod_lock_manager
|
||||
|
||||
@with_service_target(_GATEWAY_REQUEST_QUEUE_TARGET)
|
||||
async def push(self, snapshot: GatewayRequestSnapshot) -> None:
|
||||
if not snapshot:
|
||||
return
|
||||
|
|
@ -175,6 +179,7 @@ class GatewayRequestRedisBuffer:
|
|||
)
|
||||
await self._redis_cache.async_rpush(key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, values=(json.dumps(rows),))
|
||||
|
||||
@with_service_target(_GATEWAY_REQUEST_QUEUE_TARGET)
|
||||
async def _pop_batch(self) -> tuple[str | bytes, ...]:
|
||||
popped: Final[object] = await self._redis_cache.async_lpop( # pyright: ignore[reportAny] # redis returns Any
|
||||
key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT
|
||||
|
|
|
|||
|
|
@ -19,12 +19,17 @@ from datetime import datetime, timezone
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, Optional
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.proxy._types import Litellm_EntityType
|
||||
from litellm.proxy.db.db_lookup_gate import bounded_db_lookup, db_lookup_gate
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_counter, record_spend_counter_value
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import (
|
||||
SPEND_COUNTERS_TARGET,
|
||||
read_batched_spend_counter,
|
||||
record_spend_counter_value,
|
||||
)
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.project_repository import ProjectRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
|
|
@ -108,6 +113,7 @@ class SpendCounterReseed:
|
|||
return lock
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def increment_in_memory(spend_counter_cache: "DualCache", counter_key: str, increment: float) -> float | None:
|
||||
"""Apply local deltas after an in-flight reseed establishes the spend balance."""
|
||||
lock: Final = await SpendCounterReseed._get_lock(counter_key)
|
||||
|
|
@ -213,6 +219,7 @@ class SpendCounterReseed:
|
|||
return await read_batched_spend_counter(counter_key)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def coalesced(
|
||||
prisma_client: Optional["PrismaClient"],
|
||||
spend_counter_cache: "DualCache",
|
||||
|
|
@ -415,6 +422,7 @@ class SpendCounterReseed:
|
|||
return float(spend or 0.0)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def coalesced_window(
|
||||
prisma_client: Optional["PrismaClient"],
|
||||
spend_counter_cache: "DualCache",
|
||||
|
|
|
|||
|
|
@ -31,8 +31,10 @@ from fastapi import HTTPException
|
|||
|
||||
import litellm
|
||||
from litellm import DualCache
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
GUARDRAIL_SESSIONS_TARGET,
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
|
|
@ -310,6 +312,7 @@ class LassoGuardrail(CustomGuardrail):
|
|||
|
||||
return response
|
||||
|
||||
@with_service_target(GUARDRAIL_SESSIONS_TARGET)
|
||||
def _get_or_generate_conversation_id(self, data: dict, cache: DualCache) -> str:
|
||||
"""
|
||||
Get or generate a conversation_id for this request.
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import time
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy.health_check import perform_health_check
|
||||
from litellm.router_utils.health_state_cache import HEALTH_CHECKS_TARGET
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
|
@ -59,6 +61,7 @@ class SharedHealthCheckManager:
|
|||
"""Get the Redis key for model-specific health check results cache."""
|
||||
return f"health_check_results:{model_name}"
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def acquire_health_check_lock(self) -> bool:
|
||||
"""
|
||||
Attempt to acquire the global health check lock.
|
||||
|
|
@ -89,6 +92,7 @@ class SharedHealthCheckManager:
|
|||
verbose_proxy_logger.error("Error acquiring health check lock: %s", str(e))
|
||||
return False
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def release_health_check_lock(self) -> None:
|
||||
"""Release the global health check lock."""
|
||||
if self.redis_cache is None:
|
||||
|
|
@ -104,6 +108,7 @@ class SharedHealthCheckManager:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error releasing health check lock: %s", str(e))
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def get_cached_health_check_results(self) -> dict[str, Any] | None:
|
||||
"""
|
||||
Get cached health check results from Redis.
|
||||
|
|
@ -142,6 +147,7 @@ class SharedHealthCheckManager:
|
|||
verbose_proxy_logger.error("Error getting cached health check results: %s", str(e))
|
||||
return None
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def cache_health_check_results(
|
||||
self,
|
||||
healthy_endpoints: Sequence[Mapping[str, object]],
|
||||
|
|
@ -183,6 +189,7 @@ class SharedHealthCheckManager:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error caching health check results: %s", str(e))
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def perform_shared_health_check(
|
||||
self,
|
||||
model_list: list[dict[str, Any]],
|
||||
|
|
@ -319,6 +326,7 @@ class SharedHealthCheckManager:
|
|||
router=router,
|
||||
)
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def is_health_check_in_progress(self) -> bool:
|
||||
"""
|
||||
Check if a health check is currently in progress by another pod.
|
||||
|
|
@ -337,6 +345,7 @@ class SharedHealthCheckManager:
|
|||
verbose_proxy_logger.error("Error checking health check lock status: %s", str(e))
|
||||
return False
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def get_health_check_status(self) -> dict[str, object]:
|
||||
"""
|
||||
Get the current status of health check coordination.
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias
|
|||
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.constants import BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, BATCH_ENQUEUED_TOKEN_TTL_SECONDS
|
||||
|
|
@ -221,6 +222,7 @@ class BatchEnqueuedTokenStore:
|
|||
def _record_key(batch_id: str) -> str:
|
||||
return f"batch_enqueued_token_reservation:{batch_id}"
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def reserve(
|
||||
self,
|
||||
tokens: int,
|
||||
|
|
@ -325,6 +327,7 @@ class BatchEnqueuedTokenStore:
|
|||
tokens=tokens, scopes=scopes, backend="memory", owner=self._owner_token, reserved_at_monotonic=started
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def refund(
|
||||
self,
|
||||
reservation: BatchEnqueuedTokenReservation,
|
||||
|
|
@ -363,6 +366,7 @@ class BatchEnqueuedTokenStore:
|
|||
"Redis enqueued-token refund failed; leaked increments expire with the TTL: %s", str(e)
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def save_reservation(
|
||||
self,
|
||||
batch_id: str,
|
||||
|
|
@ -395,6 +399,7 @@ class BatchEnqueuedTokenStore:
|
|||
local_only=True,
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def pop_reservation(
|
||||
self,
|
||||
batch_id: str,
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from fastapi import HTTPException
|
|||
from pydantic import BaseModel, Field, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.batches.batch_utils import (
|
||||
_count_entry_tokens,
|
||||
|
|
@ -840,6 +841,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
if (descriptor := tpd_descriptors_by_counter.get(counter_key)) is not None
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def count_input_file_usage(
|
||||
self,
|
||||
file_id: str,
|
||||
|
|
@ -1177,6 +1179,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
|
||||
return file_content
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from fastapi import HTTPException
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache, InMemoryCache, RedisCache
|
||||
from litellm.caching.caching import DualCache, InMemoryCache, RedisCache, response_cache_phase
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
|
@ -63,17 +63,12 @@ class _PROXY_BatchRedisRequests(CustomLogger):
|
|||
- Get the relevant values
|
||||
"""
|
||||
if litellm.cache.type is not None and isinstance(litellm.cache.cache, RedisCache):
|
||||
# Initialize an empty list to store the keys
|
||||
keys = []
|
||||
self.print_verbose(f"cache_key_name: {cache_key_name}")
|
||||
# Use the SCAN iterator to fetch keys matching the pattern
|
||||
keys = await litellm.cache.cache.async_scan_iter(pattern=cache_key_name, count=100)
|
||||
# If you need the truly "last" based on time or another criteria,
|
||||
# ensure your key naming or storage strategy allows this determination
|
||||
# Here you would sort or filter the keys as needed based on your strategy
|
||||
self.print_verbose(f"redis keys: {keys}")
|
||||
if len(keys) > 0:
|
||||
key_value_dict = await litellm.cache.cache.async_batch_get_cache(key_list=keys)
|
||||
with response_cache_phase("get"):
|
||||
keys = await litellm.cache.cache.async_scan_iter(pattern=cache_key_name, count=100)
|
||||
self.print_verbose(f"redis keys: {keys}")
|
||||
if len(keys) > 0:
|
||||
key_value_dict = await litellm.cache.cache.async_batch_get_cache(key_list=keys)
|
||||
|
||||
## Add to cache
|
||||
if len(key_value_dict.items()) > 0:
|
||||
|
|
@ -111,7 +106,8 @@ class _PROXY_BatchRedisRequests(CustomLogger):
|
|||
max_age: Final = cache_control_args.get("s-max-age", cache_control_args.get("s-maxage", float("inf")))
|
||||
cached_result = self.in_memory_cache.get_cache(cache_key, *args, **kwargs)
|
||||
if cached_result is None:
|
||||
cached_result = await litellm.cache.cache.async_get_cache(cache_key, *args, **kwargs)
|
||||
with response_cache_phase("get"):
|
||||
cached_result = await litellm.cache.cache.async_get_cache(cache_key, *args, **kwargs)
|
||||
if cached_result is not None:
|
||||
await self.in_memory_cache.async_set_cache(cache_key, cached_result, ttl=60)
|
||||
return litellm.cache._get_cache_logic(cached_result=cached_result, max_age=max_age)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from typing import Final
|
|||
|
||||
import litellm
|
||||
from litellm import ModelResponse, Router
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.exceptions import RateLimitType
|
||||
|
|
@ -37,6 +38,7 @@ class DynamicRateLimiterCache:
|
|||
self.ttl = 60 # 1 min ttl
|
||||
self.time_fn = time_fn
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_get_cache(self, model: str) -> int | None:
|
||||
dt: Final = self.time_fn()
|
||||
current_minute: Final = dt.strftime("%H-%M")
|
||||
|
|
@ -47,6 +49,7 @@ class DynamicRateLimiterCache:
|
|||
response = len(_response)
|
||||
return response
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_set_cache_sadd(self, model: str, value: list):
|
||||
"""
|
||||
Add value to set.
|
||||
|
|
@ -82,6 +85,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
def update_variables(self, llm_router: Router):
|
||||
self.llm_router = llm_router
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def check_available_usage(
|
||||
self, model: str, priority: str | None = None
|
||||
) -> tuple[int | None, int | None, int | None, int | None, int | None]:
|
||||
|
|
@ -179,6 +183,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
)
|
||||
return None, None, None, None, None
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -234,6 +239,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
)
|
||||
return None
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
|
||||
try:
|
||||
if isinstance(response, ModelResponse):
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from fastapi import HTTPException
|
|||
|
||||
import litellm
|
||||
from litellm import ModelResponse, Router
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -569,6 +570,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
else:
|
||||
get_or_create_request_stash().rate_limit_response = atomic_response
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -656,6 +658,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
|
||||
return None
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
|
||||
"""
|
||||
Post-call hook to add rate limit headers to response.
|
||||
|
|
@ -685,6 +688,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
verbose_proxy_logger.exception("Error in dynamic rate limiter v3 post-call hook: %s", e)
|
||||
return response
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Update token usage for priority-based rate limiting after successful API calls.
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ import os
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
from litellm.exceptions import RateLimitType
|
||||
|
|
@ -83,6 +84,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
else:
|
||||
self.increment_script = None
|
||||
|
||||
@with_service_target("session_budgets")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -127,6 +129,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
|
||||
return None
|
||||
|
||||
@with_service_target("session_budgets")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
After a successful LLM call, increment the session spend by the response cost.
|
||||
|
|
@ -208,6 +211,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
def _make_cache_key(self, session_id: str) -> str:
|
||||
return f"{{session_budget:{session_id}}}:spend"
|
||||
|
||||
@with_service_target("session_budgets")
|
||||
async def _get_current_spend(self, cache_key: str) -> float:
|
||||
"""Read current accumulated spend for a session."""
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import os
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import RateLimitType
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -80,6 +81,7 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
else:
|
||||
self.increment_script = None
|
||||
|
||||
@with_service_target("session_iterations")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from typing import Final
|
|||
from openai.types import Batch
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import Span
|
||||
|
|
@ -240,6 +241,7 @@ async def build_model_max_budget_usage(
|
|||
}
|
||||
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def _current_window_spends(cache: DualCache, spend_keys: Sequence[str]) -> tuple[float, ...]:
|
||||
"""Redis holds the window total across replicas; the in-memory copy is one replica's share."""
|
||||
keys: Final = list(spend_keys)
|
||||
|
|
@ -303,6 +305,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
self._detached_increment_operations = None
|
||||
self.deployment_budget_config = None
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def is_key_within_model_budget(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -325,6 +328,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
),
|
||||
)
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def get_fallback_model_within_budget(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -339,6 +343,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
continue
|
||||
return None
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def is_user_within_model_budget(
|
||||
self,
|
||||
user_id: str,
|
||||
|
|
@ -359,6 +364,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
exceeded_message=f"LiteLLM User: {user_id}, exceeded budget for model={model}",
|
||||
)
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def is_end_user_within_model_budget(
|
||||
self,
|
||||
end_user_id: str,
|
||||
|
|
@ -379,6 +385,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
exceeded_message=f"LiteLLM End User: {end_user_id}, exceeded budget for model={model}",
|
||||
)
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def is_team_within_model_budget(
|
||||
self,
|
||||
team_id: str,
|
||||
|
|
@ -474,6 +481,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
return await self.dual_cache.async_get_cache(key=spend_key)
|
||||
return await redis_cache.async_get_cache(key=spend_key)
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def async_filter_deployments(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -484,6 +492,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
) -> list[dict]:
|
||||
return healthy_deployments
|
||||
|
||||
@with_service_target("model_budgets")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Track spend for virtual key + model in DualCache
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from typing_extensions import TypedDict
|
|||
|
||||
import litellm
|
||||
from litellm import DualCache, EmbeddingResponse, ModelResponse, TextCompletionResponse
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import RateLimitType
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -64,6 +65,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def check_key_in_limits(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -201,6 +203,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def get_all_cache_objects(
|
||||
self,
|
||||
current_global_requests: str | None,
|
||||
|
|
@ -243,6 +246,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
request_count_end_user_id=results[5],
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -489,6 +493,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
) # don't block execution for cache updates
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time):
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
get_model_group_from_litellm_kwargs,
|
||||
|
|
@ -694,6 +699,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
except Exception as e:
|
||||
self.print_verbose(e)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
self.print_verbose("Inside Max Parallel Request Failure Hook")
|
||||
|
|
@ -766,6 +772,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Inside Parallel Request Limiter: An exception occurred - %s", e)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def get_internal_user_object(
|
||||
self,
|
||||
user_id: str,
|
||||
|
|
@ -800,6 +807,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
verbose_proxy_logger.debug("Parallel Request Limiter: Error getting user object", str(e))
|
||||
return None
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
|
||||
"""
|
||||
Retrieve the key's remaining rate limits.
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from starlette.status import HTTP_503_SERVICE_UNAVAILABLE
|
|||
from typing_extensions import NotRequired, ReadOnly
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_batch import (
|
||||
BatchResult,
|
||||
|
|
@ -1169,6 +1170,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self.internal_usage_cache.dual_cache.redis_cache, RedisClusterCache
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def in_memory_cache_sliding_window(
|
||||
self,
|
||||
keys: list[str],
|
||||
|
|
@ -1525,6 +1527,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
continue
|
||||
await self._refund_counter_increments(self._counter_refunds_from_batch_values(group_keys, group_values))
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def should_rate_limit(
|
||||
self,
|
||||
descriptors: Sequence[RateLimitDescriptor],
|
||||
|
|
@ -2021,6 +2024,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
local_only=True,
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def atomic_check_and_increment_by_n(
|
||||
self,
|
||||
descriptors: list[RateLimitDescriptor],
|
||||
|
|
@ -2533,6 +2537,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
),
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def reserve_tpm_tokens(
|
||||
self,
|
||||
descriptors: list[RateLimitDescriptor],
|
||||
|
|
@ -2616,6 +2621,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def reserve_io_tokens(
|
||||
self,
|
||||
descriptors: Sequence[RateLimitDescriptor],
|
||||
|
|
@ -2701,6 +2707,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
assert itpm_response is not None
|
||||
return itpm_response, itpm_reserved, 0
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def enforce_project_io_token_quota_for_frame(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth | None,
|
||||
|
|
@ -3965,6 +3972,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
if cancellation is not None:
|
||||
raise cancellation
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -4383,6 +4391,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
batch.script(TOKEN_INCREMENT_SCRIPT, script, keys, args).on_settled(fall_back)
|
||||
return True
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_increment_tokens_with_ttl_preservation(
|
||||
self,
|
||||
pipeline_operations: list["RedisPipelineIncrementOperation"],
|
||||
|
|
@ -4494,6 +4503,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
ttl=operation["ttl"],
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_increment_reservation_aware_tokens(
|
||||
self,
|
||||
pipeline_operations: Sequence[ReservationAwareIncrementOperation],
|
||||
|
|
@ -4986,6 +4996,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
return pipeline_operations
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Update TPM usage on successful API calls by incrementing counters using pipeline
|
||||
|
|
@ -5032,6 +5043,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in rate limit success event: %s", e)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_logging_hook(
|
||||
self,
|
||||
kwargs: dict,
|
||||
|
|
@ -5102,6 +5114,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
completion_tokens,
|
||||
)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
On failure: decrement max_parallel_requests and refund the upfront
|
||||
|
|
@ -5209,6 +5222,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in rate limit failure event: %s", e)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_release_max_parallel_requests_on_disconnect(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -5229,6 +5243,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"""
|
||||
await self._release_stashed_parallel_slot(get_request_stash(), None)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
|
||||
"""
|
||||
Release completed-request slots and update rate limit headers in the response.
|
||||
|
|
@ -5283,6 +5298,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
if popped is not None:
|
||||
await self.batch_enqueued_token_store.refund(reservation=popped, litellm_parent_otel_span=span)
|
||||
|
||||
@with_service_target("rate_limits")
|
||||
async def async_post_call_failure_hook(
|
||||
self,
|
||||
request_data: dict,
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Final, Literal
|
|||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.anthropic.prompt_cache_prediction import PromptPrefix, parse_observed_cache
|
||||
|
|
@ -94,6 +95,7 @@ class PromptCacheObserver(CustomLogger):
|
|||
self.cache = internal_usage_cache.dual_cache
|
||||
self.clock = clock
|
||||
|
||||
@with_service_target("prompt_cache_predictions")
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import logging
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
|
|
@ -79,6 +80,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
|
|||
]
|
||||
return "|".join(principal) if principal else "default"
|
||||
|
||||
@with_service_target("sensitive_route_pins")
|
||||
async def _get_routed_model(self, session_id: str, user_api_key_dict: UserAPIKeyAuth | None) -> str | None:
|
||||
"""Get the model this session should be routed to, if any."""
|
||||
cache_key: Final = self._make_cache_key(session_id, self._resolve_tenant(user_api_key_dict))
|
||||
|
|
@ -114,6 +116,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
|
|||
return str(result)
|
||||
return None
|
||||
|
||||
@with_service_target("sensitive_route_pins")
|
||||
async def set_session_routing(
|
||||
self,
|
||||
session_id: str,
|
||||
|
|
@ -161,6 +164,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
|
|||
local_only=True,
|
||||
)
|
||||
|
||||
@with_service_target("sensitive_route_pins")
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from typing import Final, Protocol
|
|||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -23,6 +24,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
_get_team_object_from_cache,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_cache
|
||||
from litellm.proxy.management_helpers.resource_display_names import (
|
||||
|
|
@ -450,6 +452,7 @@ async def _patch_team_caches_remove_access_group(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _patch_key_caches_add_access_group(
|
||||
key_tokens: list[str],
|
||||
access_group_id: str,
|
||||
|
|
@ -478,6 +481,7 @@ async def _patch_key_caches_add_access_group(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _patch_key_caches_remove_access_group(
|
||||
key_tokens: list[str],
|
||||
access_group_id: str,
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from pydantic import TypeAdapter, ValidationError
|
|||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
|
@ -44,6 +45,7 @@ from litellm.proxy.auth.password_policy import (
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
object_permission_cache_key,
|
||||
user_object_permission_id_cache_key,
|
||||
)
|
||||
|
|
@ -1428,6 +1430,7 @@ def _clears_object_permission(user_request: UpdateUserRequest) -> bool:
|
|||
return sent is None or not sent.model_dump(exclude_unset=True, exclude_none=True)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _invalidate_cached_user_entitlement(user_id: str | None, object_permission_ids: tuple[str, ...]) -> None:
|
||||
"""Drop the cache entries an entitlement change makes stale.
|
||||
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from pydantic import TypeAdapter
|
|||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import service_target, with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
|
@ -82,7 +83,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import (
|
|||
)
|
||||
from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET, UserApiKeyCache
|
||||
from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks
|
||||
from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage
|
||||
from litellm.proxy.management.teams.access import TEAM_ADMIN_ONLY, TEAM_OR_ORG_ADMIN, is_team_admin
|
||||
|
|
@ -126,6 +127,7 @@ from litellm.proxy.management_helpers.team_member_permission_checks import (
|
|||
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
|
||||
from litellm.proxy.search_endpoints.search_tool_registry import rotate_search_tools_master_key
|
||||
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import SPEND_COUNTERS_TARGET
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key
|
||||
from litellm.proxy.utils import (
|
||||
PrismaClient,
|
||||
|
|
@ -3575,7 +3577,8 @@ async def update_key_fn(
|
|||
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=data.spend, ttl=60)
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=data.spend, ttl=60)
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=data.spend, ttl=60)
|
||||
except Exception as redis_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to update spend counter %s in Redis after key spend update: %s. "
|
||||
|
|
@ -4964,6 +4967,7 @@ async def can_modify_verification_token(
|
|||
return False
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def delete_verification_tokens(
|
||||
tokens: list,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
@ -6068,6 +6072,7 @@ def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_Verificatio
|
|||
return reset_to
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _set_spend_counter_with_floor_and_broadcast(counter_key: str, value: float) -> None:
|
||||
"""
|
||||
Set a Redis-backed spend counter to `value`, mirror it into the short-lived
|
||||
|
|
|
|||
|
|
@ -54,12 +54,14 @@ except ImportError:
|
|||
UniqueViolationError = Exception
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import LITELLM_PROXY_ADMIN_NAME, MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
LITELLM_MCP_SERVER_DESCRIPTION,
|
||||
LITELLM_MCP_SERVER_NAME,
|
||||
MCP_SERVERS_TARGET,
|
||||
McpServerPayloadLike,
|
||||
build_env_var_setup_url,
|
||||
collect_env_var_references,
|
||||
|
|
@ -495,6 +497,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
return server
|
||||
|
||||
@with_service_target(MCP_SERVERS_TARGET)
|
||||
async def _cache_temporary_mcp_server_in_redis(server: MCPServer, ttl_seconds: int) -> None:
|
||||
"""
|
||||
Best-effort write-through to Redis so temporary MCP OAuth sessions are
|
||||
|
|
@ -527,6 +530,7 @@ if MCP_AVAILABLE:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Failed to write temporary MCP server to Redis cache: %s", e)
|
||||
|
||||
@with_service_target(MCP_SERVERS_TARGET)
|
||||
async def _get_temporary_mcp_server_from_redis(
|
||||
server_id: str,
|
||||
) -> MCPServer | None:
|
||||
|
|
|
|||
|
|
@ -34,9 +34,11 @@ from fastapi import HTTPException, Request, status
|
|||
from fastapi.responses import RedirectResponse
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.management_endpoints.sso_helper_utils import SSO_SESSIONS_TARGET
|
||||
from litellm.proxy.management_endpoints.types import CustomOpenID, get_litellm_user_role
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
|
||||
|
|
@ -147,6 +149,7 @@ class SAMLAuthHandler:
|
|||
return SAMLAuthHandler._env("SAML_SP_ENTITY_ID") or SAMLAuthHandler._metadata_url(request)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def _load_idp_settings(cache: DualCache) -> dict[str, object]:
|
||||
metadata_url: Final = SAMLAuthHandler._env("SAML_IDP_METADATA_URL")
|
||||
metadata_xml: Final = SAMLAuthHandler._env("SAML_IDP_METADATA_XML")
|
||||
|
|
@ -241,6 +244,7 @@ class SAMLAuthHandler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def build_login_redirect(
|
||||
request: Request, cache: DualCache, relay_state: str | None = None
|
||||
) -> RedirectResponse:
|
||||
|
|
@ -358,6 +362,7 @@ class SAMLAuthHandler:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def _enforce_response_binding(
|
||||
auth: "OneLogin_Saml2_Auth",
|
||||
cache: DualCache,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
SSO_SESSIONS_TARGET: Final = "sso_sessions"
|
||||
CLI_SSO_SESSIONS_TARGET: Final = "cli_sso_sessions"
|
||||
|
||||
|
||||
def check_is_admin_only_access(ui_access_mode: str | dict) -> bool:
|
||||
"""Checks ui access mode is admin_only"""
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from fastapi.responses import RedirectResponse
|
|||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
|
@ -106,7 +107,7 @@ from litellm.proxy.common_utils.html_forms.jwt_display_template import (
|
|||
jwt_display_template,
|
||||
)
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET, UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
||||
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
|
||||
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
||||
|
|
@ -114,6 +115,8 @@ from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
|||
)
|
||||
from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
|
||||
from litellm.proxy.management_endpoints.sso_helper_utils import (
|
||||
CLI_SSO_SESSIONS_TARGET,
|
||||
SSO_SESSIONS_TARGET,
|
||||
check_is_admin_only_access,
|
||||
has_admin_ui_access,
|
||||
)
|
||||
|
|
@ -318,6 +321,7 @@ def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_fo
|
|||
return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}"
|
||||
|
||||
|
||||
@with_service_target(CLI_SSO_SESSIONS_TARGET)
|
||||
def _check_cli_sso_start_rate_limit(
|
||||
request: Request,
|
||||
cache: DualCache,
|
||||
|
|
@ -338,6 +342,7 @@ def _check_cli_sso_start_rate_limit(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(CLI_SSO_SESSIONS_TARGET)
|
||||
def _read_cli_sso_flow(cache: DualCache, cache_key: str) -> object:
|
||||
redis_cache: Final = cache.redis_cache
|
||||
if redis_cache is None:
|
||||
|
|
@ -384,6 +389,7 @@ def _get_cli_sso_flow_or_raise(login_id: str | None, cache: DualCache) -> dict:
|
|||
return flow
|
||||
|
||||
|
||||
@with_service_target(CLI_SSO_SESSIONS_TARGET)
|
||||
def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None:
|
||||
cache_key: Final = _get_cli_sso_flow_cache_key(login_id)
|
||||
redis_cache: Final = cache.redis_cache
|
||||
|
|
@ -1916,6 +1922,7 @@ def _build_sso_user_update_data(
|
|||
return update_data
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _sync_user_role_from_jwt_role_map(
|
||||
jwt_handler: JWTHandler | None,
|
||||
received_response: dict | None,
|
||||
|
|
@ -2464,6 +2471,7 @@ async def cli_sso_callback(
|
|||
|
||||
|
||||
@router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False)
|
||||
@with_service_target(CLI_SSO_SESSIONS_TARGET)
|
||||
async def cli_poll_key(
|
||||
key_id: str,
|
||||
team_id: str | None = None,
|
||||
|
|
@ -2797,6 +2805,7 @@ def _is_same_origin_return_path(return_to: str) -> bool:
|
|||
return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to)
|
||||
|
||||
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def _sso_return_to_redirect(
|
||||
return_to: str | None,
|
||||
jwt_token: str,
|
||||
|
|
@ -3060,6 +3069,7 @@ class SSOAuthenticationHandler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def get_generic_sso_redirect_response(
|
||||
generic_sso: Any,
|
||||
state: str | None = None,
|
||||
|
|
@ -3735,6 +3745,7 @@ class SSOAuthenticationHandler:
|
|||
return redirect_response
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def prepare_token_exchange_parameters(
|
||||
request: Request,
|
||||
generic_include_client_id: bool,
|
||||
|
|
@ -3914,6 +3925,7 @@ class SSOAuthenticationHandler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
@with_service_target(SSO_SESSIONS_TARGET)
|
||||
async def _delete_pkce_verifier(cache_key: str) -> None:
|
||||
"""Delete a single-use PKCE verifier from cache after a successful exchange.
|
||||
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from fastapi import HTTPException, Request
|
|||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
|
@ -36,6 +37,7 @@ from litellm.proxy._types import ( # key request types; user request types; tea
|
|||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
|
||||
from litellm.proxy.utils import PrismaClient, jsonify_object
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.table_repositories import TeamMembershipRepository
|
||||
|
|
@ -504,6 +506,7 @@ async def add_new_member(
|
|||
return returned_user, returned_team_membership
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
def _delete_user_id_from_cache(kwargs):
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
|
|
@ -518,6 +521,7 @@ def _delete_user_id_from_cache(kwargs):
|
|||
user_api_key_cache.delete_cache(key=user_id)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
def _delete_api_key_from_cache(kwargs):
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
|
|
@ -532,6 +536,7 @@ def _delete_api_key_from_cache(kwargs):
|
|||
user_api_key_cache.delete_cache(key=key)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
def _delete_team_id_from_cache(kwargs):
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
|
|
@ -546,6 +551,7 @@ def _delete_team_id_from_cache(kwargs):
|
|||
user_api_key_cache.delete_cache(key=team_id)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
def _delete_customer_id_from_cache(kwargs):
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
|
|
|
|||
|
|
@ -268,6 +268,7 @@ from functools import lru_cache, partial
|
|||
import litellm
|
||||
import litellm._redis
|
||||
from litellm import Router
|
||||
from litellm._internal_context import service_target, with_service_target
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.caching.dual_cache import DeclaredBatchRead
|
||||
|
|
@ -348,6 +349,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_team_object,
|
||||
log_db_metrics,
|
||||
)
|
||||
from litellm.proxy.auth.auth_object_prefetch import AUTH_OBJECTS_TARGET
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
check_response_size_is_safe,
|
||||
is_request_body_safe,
|
||||
|
|
@ -789,6 +791,7 @@ from litellm.proxy.spend_tracking.spend_capture_rate import (
|
|||
run_scheduled_spend_capture_rate_check,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import (
|
||||
SPEND_COUNTERS_TARGET,
|
||||
PendingSpendIncrement,
|
||||
active_spend_counter_batch,
|
||||
forget_spend_counter,
|
||||
|
|
@ -2884,6 +2887,7 @@ async def get_current_spend(
|
|||
return current
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None:
|
||||
"""Raise a counter that has fallen below the authoritative DB spend (e.g.
|
||||
Redis restarted and reloaded an older snapshot) so every worker reads the
|
||||
|
|
@ -3004,6 +3008,7 @@ async def _authoritative_floor_spend(
|
|||
return db_spend
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def read_spend_counter_cache_value(counter_key: str) -> tuple[float | None, bool]:
|
||||
"""Return (value, authoritative) for the live counter, None when absent. A clean
|
||||
Redis miss is final: the per-pod in-memory copy outlives the Redis TTL and only
|
||||
|
|
@ -3543,10 +3548,11 @@ async def _prepare_spend_counter_increment(
|
|||
under-counting (would allow overspend).
|
||||
4. Increment is returned for the caller to apply via pipeline
|
||||
"""
|
||||
await _ensure_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
)
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
await _ensure_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
)
|
||||
return PendingSpendIncrement(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
|
|
@ -3614,13 +3620,14 @@ async def _prepare_window_spend_counter_increment(
|
|||
)
|
||||
return None
|
||||
|
||||
initialized: Final = await _ensure_window_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
entity_type=entity_type,
|
||||
entity_id=entity_id,
|
||||
window_duration=window_duration,
|
||||
window_start=window_start,
|
||||
)
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
initialized: Final = await _ensure_window_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
entity_type=entity_type,
|
||||
entity_id=entity_id,
|
||||
window_duration=window_duration,
|
||||
window_start=window_start,
|
||||
)
|
||||
if initialized is False:
|
||||
return None
|
||||
return PendingSpendIncrement(counter_key=counter_key, increment=increment)
|
||||
|
|
@ -3689,6 +3696,7 @@ async def _ensure_window_spend_counter_initialized(
|
|||
return True
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _is_spend_counter_cache_warm(counter_key: str) -> bool:
|
||||
batched: Final = await read_batched_spend_counter(counter_key)
|
||||
if batched is not None:
|
||||
|
|
@ -3727,6 +3735,7 @@ async def increment_spend_counter(counter_key: str, increment: float):
|
|||
return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def refresh_spend_counter_ttl(counter_key: str) -> bool:
|
||||
if spend_counter_cache.redis_cache is None:
|
||||
return False
|
||||
|
|
@ -3737,6 +3746,7 @@ async def refresh_spend_counter_ttl(counter_key: str) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _increment_spend_counter_cache(counter_key: str, increment: float):
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
|
|
@ -3760,6 +3770,7 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float):
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def _invalidate_spend_counter(counter_key: str):
|
||||
forget_spend_counter(counter_key)
|
||||
spend_counter_cache.in_memory_cache.delete_cache(key=counter_key)
|
||||
|
|
@ -3796,8 +3807,9 @@ def _defer_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) ->
|
|||
if batch is None:
|
||||
return False
|
||||
ttl: Final = redis_cache.get_ttl()
|
||||
for item in pending:
|
||||
batch.increment(item.counter_key, item.increment, ttl).on_settled(_settle_spend_counter_increment(item))
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
for item in pending:
|
||||
batch.increment(item.counter_key, item.increment, ttl).on_settled(_settle_spend_counter_increment(item))
|
||||
return True
|
||||
|
||||
|
||||
|
|
@ -3832,6 +3844,7 @@ async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrem
|
|||
raise
|
||||
|
||||
|
||||
@with_service_target(SPEND_COUNTERS_TARGET)
|
||||
async def run_spend_counter_pipeline(pending: Sequence[PendingSpendIncrement]) -> tuple[float | None, ...]:
|
||||
"""The pipeline behind ``increment_spend_counters_pipeline`` without its invalidation: the caller decides what
|
||||
happens to counters whose increment may or may not have landed when the pipeline fails."""
|
||||
|
|
@ -3885,9 +3898,10 @@ async def arm_update_cache_read(keys: Sequence[str], cache: DualCache | None = N
|
|||
target: Final = user_api_key_cache if cache is None else cache
|
||||
if request is None or target.redis_cache is None or not keys:
|
||||
return
|
||||
request.prefetched[_UPDATE_CACHE_PREFETCH_SLOT] = await target.declare_batch_get(
|
||||
keys, request.batch(target.redis_cache)
|
||||
)
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
request.prefetched[_UPDATE_CACHE_PREFETCH_SLOT] = await target.declare_batch_get(
|
||||
keys, request.batch(target.redis_cache)
|
||||
)
|
||||
|
||||
|
||||
async def _take_armed_update_cache_read(keys: Sequence[str], cache: DualCache) -> Mapping[str, object] | None:
|
||||
|
|
@ -3945,12 +3959,13 @@ async def update_cache(
|
|||
"""
|
||||
|
||||
values_to_update_in_cache: Final[list[tuple[str, object]]] = []
|
||||
cached_values: Final = await _read_update_cache_values(
|
||||
keys=update_cache_read_keys(
|
||||
user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost
|
||||
),
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
cached_values: Final = await _read_update_cache_values(
|
||||
keys=update_cache_read_keys(
|
||||
user_id=user_id, end_user_id=end_user_id, team_id=team_id, tags=tags, response_cost=response_cost
|
||||
),
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
### UPDATE KEY SPEND ###
|
||||
async def _update_key_cache(token: str, response_cost: float):
|
||||
|
|
@ -4194,42 +4209,45 @@ async def update_cache(
|
|||
traceback.format_exc(),
|
||||
)
|
||||
|
||||
if token is not None and response_cost is not None:
|
||||
await _update_key_cache(token=token, response_cost=response_cost)
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
if token is not None and response_cost is not None:
|
||||
await _update_key_cache(token=token, response_cost=response_cost)
|
||||
|
||||
if user_id is not None:
|
||||
await _update_user_cache()
|
||||
if user_id is not None:
|
||||
await _update_user_cache()
|
||||
|
||||
if end_user_id is not None:
|
||||
await _update_end_user_cache()
|
||||
if end_user_id is not None:
|
||||
await _update_end_user_cache()
|
||||
|
||||
if team_id is not None:
|
||||
await _update_team_cache()
|
||||
if team_id is not None:
|
||||
await _update_team_cache()
|
||||
|
||||
if tags is not None:
|
||||
await _update_tag_cache()
|
||||
if tags is not None:
|
||||
await _update_tag_cache()
|
||||
|
||||
global_proxy_spend_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
|
||||
local_object_updates: Final = tuple((k, v) for k, v in values_to_update_in_cache if k != global_proxy_spend_key)
|
||||
shared_scalar_updates: Final = tuple((k, v) for k, v in values_to_update_in_cache if k == global_proxy_spend_key)
|
||||
|
||||
if local_object_updates:
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=list(local_object_updates),
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
local_only=True,
|
||||
with service_target(AUTH_OBJECTS_TARGET):
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=list(local_object_updates),
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
local_only=True,
|
||||
)
|
||||
)
|
||||
)
|
||||
if shared_scalar_updates:
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=list(shared_scalar_updates),
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=list(shared_scalar_updates),
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def run_ollama_serve():
|
||||
|
|
|
|||
|
|
@ -6,11 +6,14 @@ import json
|
|||
from datetime import datetime, timezone
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid4
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesAPIStatus
|
||||
|
||||
_RESPONSE_POLLING_TARGET: Final = "response_polling"
|
||||
|
||||
|
||||
class ResponsePollingHandler:
|
||||
"""Handles polling-based responses with Redis cache"""
|
||||
|
|
@ -37,6 +40,7 @@ class ResponsePollingHandler:
|
|||
"""Get Redis cache key for a polling ID"""
|
||||
return f"{cls.CACHE_KEY_PREFIX}{polling_id}"
|
||||
|
||||
@with_service_target(_RESPONSE_POLLING_TARGET)
|
||||
async def create_initial_state(
|
||||
self,
|
||||
polling_id: str,
|
||||
|
|
@ -81,6 +85,7 @@ class ResponsePollingHandler:
|
|||
|
||||
return response
|
||||
|
||||
@with_service_target(_RESPONSE_POLLING_TARGET)
|
||||
async def update_state(
|
||||
self,
|
||||
polling_id: str,
|
||||
|
|
@ -212,6 +217,7 @@ class ResponsePollingHandler:
|
|||
"Updated polling state for %s: status=%s, output_items=%s", polling_id, state["status"], output_count
|
||||
)
|
||||
|
||||
@with_service_target(_RESPONSE_POLLING_TARGET)
|
||||
async def get_state(self, polling_id: str) -> dict[str, Any] | None:
|
||||
"""Get current polling state from Redis"""
|
||||
if not self.redis_cache:
|
||||
|
|
@ -237,6 +243,7 @@ class ResponsePollingHandler:
|
|||
)
|
||||
return True
|
||||
|
||||
@with_service_target(_RESPONSE_POLLING_TARGET)
|
||||
async def delete_polling(self, polling_id: str) -> bool:
|
||||
"""Delete a polling request from cache"""
|
||||
if not self.redis_cache:
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import Final, NoReturn, SupportsFloat, SupportsIndex, SupportsInt, c
|
|||
from fastapi import HTTPException, status
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import select_tier_for_input, tier_rate
|
||||
|
|
@ -27,6 +28,7 @@ from litellm.proxy.auth.auth_utils import get_model_from_request
|
|||
from litellm.proxy.auth.budget_throttle import should_throttle_budget_exceeded
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
AUTH_OBJECTS_TARGET,
|
||||
UserApiKeyCache,
|
||||
end_user_cache_key,
|
||||
model_access_group_cache_key,
|
||||
|
|
@ -732,6 +734,7 @@ def _dedupe_tags(tags: list[str]) -> list[str]:
|
|||
return deduped_tags
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _get_team_member_budget_counter(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
|
|
@ -782,6 +785,7 @@ async def _get_team_member_budget_counter(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _get_org_budget_counter(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
|
|
@ -820,6 +824,7 @@ async def _get_org_budget_counter(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(AUTH_OBJECTS_TARGET)
|
||||
async def _get_project_budget_counter(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from datetime import date, datetime, time, timedelta, timezone
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
PTU_LAPSED_ALERT_LIMIT,
|
||||
|
|
@ -30,6 +31,7 @@ from litellm.constants import (
|
|||
PTU_SENTINEL_API_KEY,
|
||||
)
|
||||
from litellm.litellm_core_utils.ptu_pricing import ptu_terms
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import POD_LOCK_TARGET
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
|
||||
from litellm.repositories.model_repository import ModelRepository
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
|
|
@ -651,6 +653,7 @@ async def run_scheduled_ptu_rollup(
|
|||
await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID)
|
||||
|
||||
|
||||
@with_service_target(POD_LOCK_TARGET)
|
||||
async def _lock_is_held(pod_lock_manager: "PodLockManager") -> bool:
|
||||
"""True only when the rollup lock is readable and someone is holding it.
|
||||
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from typing import TYPE_CHECKING, Final, TypeAlias
|
|||
from pydantic import BaseModel, ConfigDict, TypeAdapter
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
|
||||
|
|
@ -27,6 +28,7 @@ from litellm.llms.openai.organization_costs import (
|
|||
fetch_openai_daily_costs,
|
||||
provider_billing_get,
|
||||
)
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import POD_LOCK_TARGET
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.proxy.spend_capture_rate import (
|
||||
CaptureRateDay,
|
||||
|
|
@ -290,6 +292,7 @@ async def _claims_alert_window(pod_lock_manager: "PodLockManager | None") -> boo
|
|||
return acquired or not await _lock_is_held(pod_lock_manager, redis_cache)
|
||||
|
||||
|
||||
@with_service_target(POD_LOCK_TARGET)
|
||||
async def _lock_is_held(pod_lock_manager: "PodLockManager", redis_cache: "RedisCache") -> bool:
|
||||
try:
|
||||
return bool(
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from typing import Final
|
|||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.redis_batch import BatchResult, RedisBatch, active_request_redis_batch
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
|
@ -20,6 +21,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
|
||||
_CounterValues: Final = TypeAdapter(dict[str, float | None])
|
||||
_NO_VALUES: Final[Mapping[str, float | None]] = MappingProxyType({})
|
||||
SPEND_COUNTERS_TARGET: Final = "spend_counters"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -112,7 +114,8 @@ class SpendCounterBatch:
|
|||
pending: Final = self._keys - self._fetched
|
||||
if pending:
|
||||
self._fetched = self._fetched | pending
|
||||
self._inflight.append(self._request_batch.mget(sorted(pending)))
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
self._inflight.append(self._request_batch.mget(sorted(pending)))
|
||||
|
||||
async def _collect_inflight(self) -> None:
|
||||
results: Final = tuple(self._inflight)
|
||||
|
|
@ -127,9 +130,9 @@ class SpendCounterBatch:
|
|||
|
||||
async def _fetch(self, keys: frozenset[str]) -> Mapping[str, float | None]:
|
||||
try:
|
||||
return _CounterValues.validate_python(
|
||||
await self._redis_cache.async_batch_get_cache(key_list=sorted(keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped cache API
|
||||
)
|
||||
with service_target(SPEND_COUNTERS_TARGET):
|
||||
values: Final = await self._redis_cache.async_batch_get_cache(key_list=sorted(keys)) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped cache API
|
||||
return _CounterValues.validate_python(values) # pyright: ignore[reportUnknownArgumentType] # untyped cache API
|
||||
except Exception as e: # noqa: BLE001 # per-key reads take over and apply their own Redis fallback
|
||||
verbose_proxy_logger.debug("spend counter batch read failed, falling back to per-key reads: %s", e)
|
||||
return _NO_VALUES
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from pydantic.fields import FieldInfo, PydanticUndefined
|
|||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import MCP_TOOL_SEARCH_SETTINGS_KEY
|
||||
|
|
@ -41,7 +42,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
|||
PTU_COST_ATTRIBUTION_ENV_VAR,
|
||||
is_ptu_cost_attribution_enabled,
|
||||
)
|
||||
from litellm.proxy.utils import invalidate_config_param
|
||||
from litellm.proxy.utils import CONFIG_PARAMS_TARGET, invalidate_config_param
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
|
|
@ -1668,6 +1669,7 @@ UI_SETTINGS_CACHE_KEY: Final = "ui_settings:settings_dict"
|
|||
UI_SETTINGS_CACHE_TTL: Final = 600 # 10 minutes
|
||||
|
||||
|
||||
@with_service_target(CONFIG_PARAMS_TARGET)
|
||||
async def get_ui_settings_cached() -> dict[str, JsonValue]:
|
||||
"""
|
||||
Return the persisted UI settings dict, using DualCache for reads.
|
||||
|
|
@ -1747,6 +1749,7 @@ async def sync_ui_settings_to_general_settings(prisma_client: object) -> Mapping
|
|||
tags=["UI Settings"],
|
||||
response_model=UISettingsResponse,
|
||||
)
|
||||
@with_service_target(CONFIG_PARAMS_TARGET)
|
||||
async def get_ui_settings():
|
||||
"""
|
||||
Get UI-specific configuration flags.
|
||||
|
|
@ -1825,6 +1828,7 @@ async def get_ui_settings():
|
|||
tags=["UI Settings"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@with_service_target(CONFIG_PARAMS_TARGET)
|
||||
async def update_ui_settings(
|
||||
settings_body: dict[str, object] = Body(...),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
|
|
|
|||
|
|
@ -119,6 +119,7 @@ from litellm import (
|
|||
ModelResponseStream,
|
||||
Router,
|
||||
)
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging, ServiceTypes
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
|
|
@ -4268,6 +4269,9 @@ class _ConfigRow:
|
|||
self.param_value = param_value
|
||||
|
||||
|
||||
CONFIG_PARAMS_TARGET: Final = "config_params"
|
||||
|
||||
|
||||
def _config_cache_key(param_name: str) -> str:
|
||||
return f"litellm_config:param:{param_name}"
|
||||
|
||||
|
|
@ -4287,18 +4291,21 @@ def _unpack_config_row(cached: object) -> _ConfigRow | None:
|
|||
async def get_config_param(prisma_client: "PrismaClient", param_name: str) -> Any | None:
|
||||
"""Cached read of a LiteLLM_Config row; returns row, _ConfigRow shim, or None."""
|
||||
cache_key: Final = _config_cache_key(param_name)
|
||||
cached: Final = await litellm_config_cache.async_get_cache(cache_key)
|
||||
with service_target(CONFIG_PARAMS_TARGET):
|
||||
cached: Final = await litellm_config_cache.async_get_cache(cache_key)
|
||||
if cached is not None:
|
||||
return _unpack_config_row(cached)
|
||||
|
||||
row: Final = await prisma_client.get_generic_data(key="param_name", value=param_name, table_name="config")
|
||||
cache_value: Final[Mapping[str, object] | str] = _pack_config_row(row) if row is not None else _CONFIG_CACHE_MISS
|
||||
await litellm_config_cache.async_set_cache(cache_key, cache_value, ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS)
|
||||
with service_target(CONFIG_PARAMS_TARGET):
|
||||
await litellm_config_cache.async_set_cache(cache_key, cache_value, ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS)
|
||||
return row
|
||||
|
||||
|
||||
async def evict_config_param(param_name: str) -> None:
|
||||
await litellm_config_cache.async_delete_cache(_config_cache_key(param_name))
|
||||
with service_target(CONFIG_PARAMS_TARGET):
|
||||
await litellm_config_cache.async_delete_cache(_config_cache_key(param_name))
|
||||
|
||||
|
||||
async def invalidate_config_param(param_name: str) -> None:
|
||||
|
|
@ -4323,12 +4330,13 @@ async def prefetch_config_params(prisma_client: "PrismaClient | None", param_nam
|
|||
)
|
||||
return
|
||||
by_name: Final = {row.param_name: row for row in rows}
|
||||
for name in param_names:
|
||||
row = by_name.get(name)
|
||||
cache_value: Mapping[str, object] | str = _pack_config_row(row) if row is not None else _CONFIG_CACHE_MISS
|
||||
await litellm_config_cache.async_set_cache(
|
||||
_config_cache_key(name), cache_value, ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS
|
||||
)
|
||||
with service_target(CONFIG_PARAMS_TARGET):
|
||||
for name in param_names:
|
||||
row = by_name.get(name)
|
||||
cache_value: Mapping[str, object] | str = _pack_config_row(row) if row is not None else _CONFIG_CACHE_MISS
|
||||
await litellm_config_cache.async_set_cache(
|
||||
_config_cache_key(name), cache_value, ttl=LITELLM_CONFIG_CACHE_TTL_SECONDS
|
||||
)
|
||||
|
||||
|
||||
_WRITER_WRITABILITY_PROBE_SQL: Final = "SELECT current_setting('transaction_read_only') AS transaction_read_only"
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ from typing_extensions import overload
|
|||
import litellm
|
||||
import litellm.litellm_core_utils.exception_mapping_utils
|
||||
from litellm import get_secret_str
|
||||
from litellm._internal_context import service_target, with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.caching import (
|
||||
|
|
@ -71,6 +72,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.integrations.custom_guardrail import is_guardrail_intervention
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.otel.runtime import phase_span
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
|
|
@ -260,7 +262,7 @@ from litellm.router_utils.routing_groups import (
|
|||
parse_routing_groups,
|
||||
validate_routing_strategy,
|
||||
)
|
||||
from litellm.router_utils.routing_read_batch import RoutingPrefetch, RoutingReadBatch
|
||||
from litellm.router_utils.routing_read_batch import ROUTER_USAGE_TARGET, RoutingPrefetch, RoutingReadBatch
|
||||
from litellm.scheduler import FlowItem, Scheduler
|
||||
from litellm.types.litellm_params import RoutingStrategyName
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -441,6 +443,7 @@ _ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key"
|
|||
_ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params"
|
||||
_CLAUDE_CODE_SESSION_ID_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
|
||||
_CLAUDE_CODE_SESSION_ROUTER_TTL_SECONDS: Final = 3600
|
||||
CLAUDE_CODE_SESSION_ROUTER_BINDING_TARGET: Final = "claude_code_session_router_binding"
|
||||
|
||||
_RUNTIME_TOGGLEABLE_PRE_CALL_CHECKS: Final[Mapping[str, type[CustomLogger]]] = MappingProxyType(
|
||||
{
|
||||
|
|
@ -8347,12 +8350,13 @@ class Router:
|
|||
|
||||
## RPM
|
||||
rpm_key: Final = RouterCacheEnum.RPM.value.format(id=id, current_minute=current_minute, model=deployment_name)
|
||||
await self.cache.async_increment_cache(
|
||||
key=rpm_key,
|
||||
value=1,
|
||||
parent_otel_span=parent_otel_span,
|
||||
ttl=RoutingArgs.ttl.value,
|
||||
)
|
||||
with service_target(ROUTER_USAGE_TARGET):
|
||||
await self.cache.async_increment_cache(
|
||||
key=rpm_key,
|
||||
value=1,
|
||||
parent_otel_span=parent_otel_span,
|
||||
ttl=RoutingArgs.ttl.value,
|
||||
)
|
||||
|
||||
def _get_metadata_variable_name_from_kwargs(self, kwargs: dict) -> Literal["metadata", "litellm_metadata"]:
|
||||
"""
|
||||
|
|
@ -13274,6 +13278,23 @@ class Router:
|
|||
|
||||
Allows all cache calls to be made async => 10x perf impact (8rps -> 100 rps).
|
||||
"""
|
||||
with phase_span(f"route {model}"):
|
||||
return await self._async_get_available_deployment(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
async def _async_get_available_deployment(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: dict,
|
||||
messages: list[dict[str, str]] | None,
|
||||
input: str | list | None,
|
||||
specific_deployment: bool | None,
|
||||
):
|
||||
if (
|
||||
self.routing_strategy != "usage-based-routing-v2"
|
||||
and self.routing_strategy != "simple-shuffle"
|
||||
|
|
@ -13425,6 +13446,23 @@ class Router:
|
|||
|
||||
Only returns deployments configured with use_in_pass_through=True
|
||||
"""
|
||||
with phase_span(f"route {model}"):
|
||||
return await self._async_get_available_deployment_for_pass_through(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
async def _async_get_available_deployment_for_pass_through(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: dict,
|
||||
messages: list[dict[str, str]] | None,
|
||||
input: str | list | None,
|
||||
specific_deployment: bool | None,
|
||||
):
|
||||
try:
|
||||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs)
|
||||
|
||||
|
|
@ -13709,6 +13747,7 @@ class Router:
|
|||
return None
|
||||
return f"claude_code_session_router:v1:{caller_scope}:{session_id}"
|
||||
|
||||
@with_service_target(CLAUDE_CODE_SESSION_ROUTER_BINDING_TARGET)
|
||||
async def _delete_claude_code_session_router_binding(self, cache_key: str) -> None:
|
||||
try:
|
||||
await self._claude_code_session_router_cache.async_delete_cache(key=cache_key)
|
||||
|
|
@ -13718,6 +13757,7 @@ class Router:
|
|||
e,
|
||||
)
|
||||
|
||||
@with_service_target(CLAUDE_CODE_SESSION_ROUTER_BINDING_TARGET)
|
||||
async def _get_claude_code_session_router_binding(self, cache_key: str) -> object:
|
||||
session_cache: Final = self._claude_code_session_router_cache
|
||||
try:
|
||||
|
|
@ -13733,6 +13773,7 @@ class Router:
|
|||
)
|
||||
return None
|
||||
|
||||
@with_service_target(CLAUDE_CODE_SESSION_ROUTER_BINDING_TARGET)
|
||||
async def _resolve_claude_code_session_router(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import logging
|
|||
from abc import ABC
|
||||
from typing import Final
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisPipelineIncrementOperation, log_redis_failure
|
||||
|
|
@ -99,6 +100,7 @@ class BaseRoutingStrategy(ABC):
|
|||
self.add_to_in_memory_keys_to_update(key=key)
|
||||
return result
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def periodic_sync_in_memory_spend_with_redis(self, default_sync_interval: float | None):
|
||||
"""
|
||||
Handler that triggers sync_in_memory_spend_with_redis every DEFAULT_REDIS_SYNC_INTERVAL seconds
|
||||
|
|
@ -118,6 +120,7 @@ class BaseRoutingStrategy(ABC):
|
|||
default_sync_interval
|
||||
) # Still wait DEFAULT_REDIS_SYNC_INTERVAL seconds on error before retrying
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def _push_in_memory_increments_to_redis(self):
|
||||
"""
|
||||
How this works:
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from types import MappingProxyType
|
|||
from typing import Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCache, RedisPipelineIncrementOperation, log_redis_failure
|
||||
|
|
@ -127,6 +128,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
if isinstance(litellm.callbacks, list):
|
||||
litellm.logging_callback_manager.add_litellm_callback(self)
|
||||
|
||||
@with_service_target("router_budgets")
|
||||
async def async_filter_deployments(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -468,6 +470,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
flush_task.result()
|
||||
raise
|
||||
|
||||
@with_service_target("router_budgets")
|
||||
async def _write_queued_increment_operations(self, redis_cache: RedisCache) -> bool:
|
||||
increment_operations_to_flush: Final = await self._detach_queued_increment_operations()
|
||||
if len(increment_operations_to_flush) == 0:
|
||||
|
|
@ -488,6 +491,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
await self._clear_detached_increment_operations()
|
||||
return True
|
||||
|
||||
@with_service_target("router_budgets")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""Original method now uses helper functions"""
|
||||
verbose_router_logger.debug("in RouterBudgetLimiting.async_log_success_event")
|
||||
|
|
@ -594,6 +598,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
|
||||
verbose_router_logger.debug("Incremented spend for %s by %s", spend_key, response_cost)
|
||||
|
||||
@with_service_target("router_budgets")
|
||||
async def periodic_sync_in_memory_spend_with_redis(self):
|
||||
"""
|
||||
Handler that triggers sync_in_memory_spend_with_redis every DEFAULT_REDIS_SYNC_INTERVAL seconds
|
||||
|
|
@ -750,6 +755,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
budget_limit=budget_limit,
|
||||
)
|
||||
|
||||
@with_service_target("router_budgets")
|
||||
async def _get_current_provider_spend(self, provider: str) -> float | None:
|
||||
"""
|
||||
GET the current spend for a provider from cache
|
||||
|
|
@ -776,6 +782,7 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
current_spend = await self.dual_cache.async_get_cache(spend_key)
|
||||
return float(current_spend) if current_spend is not None else 0.0
|
||||
|
||||
@with_service_target("router_budgets")
|
||||
async def _get_current_provider_budget_reset_at(self, provider: str) -> str | None:
|
||||
budget_config: Final = self._get_budget_config_for_provider(provider)
|
||||
if budget_config is None:
|
||||
|
|
|
|||
|
|
@ -31,8 +31,9 @@ from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
|
|||
from pydantic import BaseModel, TypeAdapter, ValidationError, create_model
|
||||
from pydantic_core import ErrorDetails
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.affinity_cache import claim_affinity_pin
|
||||
from litellm.caching.affinity_cache import ROUTER_SESSION_PINS_TARGET, claim_affinity_pin
|
||||
from litellm.constants import (
|
||||
EMPTY_MAPPING,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
|
|
@ -4180,6 +4181,7 @@ class ComplexityRouter(CustomLogger):
|
|||
return response
|
||||
return response.model_copy(update={"session_affinity_ttl_seconds": self.config.session_affinity_ttl_seconds})
|
||||
|
||||
@with_service_target(ROUTER_SESSION_PINS_TARGET)
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from typing import Final
|
|||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import log_redis_failure
|
||||
|
|
@ -119,9 +120,11 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
self.router_cache = router_cache
|
||||
self.router_cache_id = str(id(router_cache))
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
|
||||
self._increment(kwargs, 1)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
|
|
@ -129,6 +132,7 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
|
|
@ -136,6 +140,7 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
|
|
@ -143,6 +148,7 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
if self.test_flag:
|
||||
self.logged_success += 1
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
|
|
@ -150,6 +156,7 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
if self.test_flag:
|
||||
self.logged_failure += 1
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def get_available_deployments(
|
||||
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
|
||||
) -> Mapping[str, object] | None:
|
||||
|
|
@ -165,6 +172,7 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
local: Final = _local_counts(self.router_cache.batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_get_available_deployments(
|
||||
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
|
||||
) -> Mapping[str, object] | None:
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from typing import Final
|
|||
|
||||
import litellm
|
||||
from litellm import ModelResponse, token_counter, verbose_logger
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -19,6 +20,7 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
def __init__(self, router_cache: DualCache, routing_args: dict = {}):
|
||||
self.router_cache = router_cache
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -94,6 +96,7 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
"litellm.router_strategy.lowest_cost.py::log_success_event(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -169,6 +172,7 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from pydantic import Field
|
|||
|
||||
import litellm
|
||||
from litellm import ModelResponse, token_counter, verbose_logger
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, safe_divide_seconds
|
||||
|
|
@ -58,6 +59,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
self.router_cache = router_cache
|
||||
self.routing_args = RoutingArgs(**routing_args)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -181,6 +183,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Check if Timeout Error, if timeout set deployment latency -> 100
|
||||
|
|
@ -240,6 +243,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
"litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -497,6 +501,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
request_kwargs[metadata_field]["_latency_per_deployment"] = _latency_per_deployment
|
||||
return deployment
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
|
|
@ -522,6 +527,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
request_count_dict,
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from datetime import datetime
|
|||
from typing import Final
|
||||
|
||||
from litellm import token_counter
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -27,6 +28,7 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
self.router_cache = router_cache
|
||||
self.routing_args = RoutingArgs(**routing_args)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -81,6 +83,7 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
)
|
||||
verbose_router_logger.debug(traceback.format_exc())
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -145,6 +148,7 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
)
|
||||
verbose_router_logger.debug(traceback.format_exc())
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import httpx
|
|||
|
||||
import litellm
|
||||
from litellm import token_counter
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger, verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -98,6 +99,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
default_sync_interval=0.1,
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def pre_call_check(self, deployment: dict) -> dict | None:
|
||||
"""
|
||||
Pre-call check + update model rpm
|
||||
|
|
@ -173,6 +175,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
raise e
|
||||
return deployment # don't fail calls if eg. redis fails to connect
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_pre_call_check(self, deployment: dict, parent_otel_span: Span | None) -> dict | None:
|
||||
"""
|
||||
Pre-call check + update model rpm
|
||||
|
|
@ -249,6 +252,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
raise e
|
||||
return deployment # don't fail calls if eg. redis fails to connect
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -291,6 +295,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
"litellm.proxy.hooks.lowest_tpm_rpm_v2.py::log_success_event(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
|
|
@ -464,6 +469,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
[f"{prefix}:rpm:{current_minute}" for prefix in prefixes],
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
async def async_get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
|
|
@ -572,6 +578,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
|
|||
),
|
||||
)
|
||||
|
||||
@with_service_target("router_usage")
|
||||
def get_available_deployments(
|
||||
self,
|
||||
model_group: str,
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._internal_context import service_target
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import DEFAULT_COOLDOWN_REDIS_READ_INTERVAL_SECONDS
|
||||
|
|
@ -34,6 +35,7 @@ class CooldownCacheValue(TypedDict):
|
|||
# real remaining cooldown against Redis at least this often, so an entry that later gets
|
||||
# deleted or extended in Redis before its original deadline is still noticed promptly.
|
||||
_MAX_CORRECTED_IN_MEMORY_TTL_SECONDS: Final = 60.0
|
||||
ROUTER_COOLDOWNS_TARGET: Final = "router_cooldowns"
|
||||
|
||||
|
||||
class CooldownCache:
|
||||
|
|
@ -118,11 +120,12 @@ class CooldownCache:
|
|||
)
|
||||
|
||||
# Set the cache with a TTL equal to the cooldown time
|
||||
self.cooldown_store.set_cache(
|
||||
value=cooldown_data,
|
||||
key=cooldown_key,
|
||||
ttl=_cooldown_time,
|
||||
)
|
||||
with service_target(ROUTER_COOLDOWNS_TARGET):
|
||||
self.cooldown_store.set_cache(
|
||||
value=cooldown_data,
|
||||
key=cooldown_key,
|
||||
ttl=_cooldown_time,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error("CooldownCache::add_deployment_to_cooldown - Exception occurred - %s", e)
|
||||
raise e
|
||||
|
|
@ -162,7 +165,10 @@ class CooldownCache:
|
|||
# Generate the keys for the deployments
|
||||
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
|
||||
results: Final = await self.cooldown_store.async_batch_get_cache(keys=keys, parent_otel_span=parent_otel_span)
|
||||
with service_target(ROUTER_COOLDOWNS_TARGET):
|
||||
results: Final = await self.cooldown_store.async_batch_get_cache(
|
||||
keys=keys, parent_otel_span=parent_otel_span
|
||||
)
|
||||
return self.active_cooldowns_from_results(model_ids, results)
|
||||
|
||||
def active_cooldowns_from_results(
|
||||
|
|
@ -190,7 +196,8 @@ class CooldownCache:
|
|||
# Generate the keys for the deployments
|
||||
keys: Final = [CooldownCache.get_cooldown_cache_key(model_id) for model_id in model_ids]
|
||||
# Retrieve the values for the keys using mget
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
with service_target(ROUTER_COOLDOWNS_TARGET):
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
active_cooldowns: Final = []
|
||||
current_time: Final = time.time()
|
||||
|
|
@ -210,7 +217,8 @@ class CooldownCache:
|
|||
keys: Final = [f"deployment:{model_id}:cooldown" for model_id in model_ids]
|
||||
|
||||
# Retrieve the values for the keys using mget
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
with service_target(ROUTER_COOLDOWNS_TARGET):
|
||||
results: Final = self.cooldown_store.batch_get_cache(keys=keys, parent_otel_span=parent_otel_span) or []
|
||||
|
||||
min_cooldown_time: float | None = None
|
||||
# Process the results
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -23,6 +24,7 @@ from litellm.constants import (
|
|||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD,
|
||||
)
|
||||
from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET
|
||||
from litellm.router_utils.cooldown_callbacks import router_cooldown_event_callback
|
||||
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
|
||||
|
||||
|
|
@ -614,12 +616,13 @@ def _increment_allowed_fails(cache: DualCache, cache_key: str, ttl: float) -> in
|
|||
Return the fleet-wide fail count. ``DualCache.increment_cache`` bumps the in-memory tier
|
||||
before Redis and re-raises a Redis error, so a Redis outage degrades to this worker's own count.
|
||||
"""
|
||||
try:
|
||||
return cache.increment_cache(key=cache_key, value=1, ttl=ttl)
|
||||
except Exception as e: # noqa: BLE001 # a Redis outage must not stop failing deployments from cooling down
|
||||
verbose_router_logger.warning("allowed_fails counter fell back to this worker's in-memory count: %s", e)
|
||||
local_fails: Final = cache.get_cache(key=cache_key, local_only=True)
|
||||
return local_fails if isinstance(local_fails, int) else 0
|
||||
with service_target(ROUTER_COOLDOWNS_TARGET):
|
||||
try:
|
||||
return cache.increment_cache(key=cache_key, value=1, ttl=ttl)
|
||||
except Exception as e: # noqa: BLE001 # a Redis outage must not stop failing deployments from cooling down
|
||||
verbose_router_logger.warning("allowed_fails counter fell back to this worker's in-memory count: %s", e)
|
||||
local_fails: Final = cache.get_cache(key=cache_key, local_only=True)
|
||||
return local_fails if isinstance(local_fails, int) else 0
|
||||
|
||||
|
||||
def _is_allowed_fails_set_on_router(
|
||||
|
|
|
|||
|
|
@ -11,9 +11,12 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
|
||||
HEALTH_CHECKS_TARGET: Final = "health_checks"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
|
|
@ -28,6 +31,7 @@ class DeploymentHealthStateValue(TypedDict):
|
|||
reason: str
|
||||
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
def _read_shared_health_snapshot(cache: DualCache, key: str) -> object:
|
||||
redis_cache: Final = cache.redis_cache
|
||||
if redis_cache is None:
|
||||
|
|
@ -53,6 +57,7 @@ class DeploymentHealthCache:
|
|||
self.cache = cache
|
||||
self.staleness_threshold = staleness_threshold
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
def set_deployment_health_states(self, states: dict[str, DeploymentHealthStateValue]) -> None:
|
||||
"""Merge the given states into the shared cache entry, pruning expired ones.
|
||||
|
||||
|
|
@ -100,6 +105,7 @@ class DeploymentHealthCache:
|
|||
and (now - state.get("timestamp", 0)) < self.staleness_threshold
|
||||
}
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
async def async_get_unhealthy_deployment_ids(self, parent_otel_span: Span | None = None) -> set[str]:
|
||||
"""Return set of deployment IDs currently marked unhealthy and not stale."""
|
||||
try:
|
||||
|
|
@ -112,6 +118,7 @@ class DeploymentHealthCache:
|
|||
)
|
||||
return set()
|
||||
|
||||
@with_service_target(HEALTH_CHECKS_TARGET)
|
||||
def get_unhealthy_deployment_ids(self, parent_otel_span: Span | None = None) -> set[str]:
|
||||
"""Sync version: return set of deployment IDs currently marked unhealthy and not stale."""
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -18,8 +18,14 @@ from typing import Any, Final, cast
|
|||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.affinity_cache import claim_affinity_pin, claim_affinity_pin_in_memory, set_local_affinity_pin
|
||||
from litellm.caching.affinity_cache import (
|
||||
ROUTER_SESSION_PINS_TARGET,
|
||||
claim_affinity_pin,
|
||||
claim_affinity_pin_in_memory,
|
||||
set_local_affinity_pin,
|
||||
)
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, SESSION_ID_GENERATED_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
|
|
@ -345,6 +351,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return deployment
|
||||
return None
|
||||
|
||||
@with_service_target(ROUTER_SESSION_PINS_TARGET)
|
||||
async def async_filter_deployments(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -19,9 +19,11 @@ import httpx
|
|||
|
||||
import litellm
|
||||
from litellm import token_counter
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.litellm_core_utils.token_counter import offload_token_count
|
||||
from litellm.router_utils.routing_read_batch import ROUTER_USAGE_TARGET
|
||||
from litellm.types.router import RouterCacheEnum, RouterErrors
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
||||
|
|
@ -343,6 +345,7 @@ def _rate_limit_error(limit_label: str, limit: int, current: float) -> litellm.R
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def _sync_increment_with_rollback(
|
||||
dual_cache: DualCache,
|
||||
key: str,
|
||||
|
|
@ -367,6 +370,7 @@ def _sync_increment_with_rollback(
|
|||
raise _rate_limit_error(limit_label, limit, current)
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def _increment_with_rollback(
|
||||
dual_cache: DualCache,
|
||||
key: str,
|
||||
|
|
@ -394,6 +398,7 @@ async def _increment_with_rollback(
|
|||
raise _rate_limit_error(limit_label, limit, current)
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def io_token_pre_call_check(
|
||||
dual_cache: DualCache,
|
||||
deployment: dict,
|
||||
|
|
@ -456,6 +461,7 @@ def io_token_pre_call_check(
|
|||
return deployment
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def async_io_token_pre_call_check(
|
||||
dual_cache: DualCache,
|
||||
deployment: dict,
|
||||
|
|
@ -525,6 +531,7 @@ async def async_io_token_pre_call_check(
|
|||
return deployment
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def io_token_reconcile_success(
|
||||
dual_cache: DualCache,
|
||||
kwargs: Mapping[str, object] | None,
|
||||
|
|
@ -576,6 +583,7 @@ def io_token_reconcile_success(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def async_io_token_reconcile_success(
|
||||
dual_cache: DualCache,
|
||||
kwargs: Mapping[str, object] | None,
|
||||
|
|
@ -637,6 +645,7 @@ async def async_io_token_reconcile_success(
|
|||
)
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def io_token_refund_failure(
|
||||
dual_cache: DualCache,
|
||||
kwargs: Mapping[str, object] | None,
|
||||
|
|
@ -688,6 +697,7 @@ def refund_stale_reservation_before_retry(dual_cache: DualCache, kwargs: Mapping
|
|||
io_token_refund_failure(dual_cache, kwargs)
|
||||
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def async_io_token_refund_failure(
|
||||
dual_cache: DualCache,
|
||||
kwargs: Mapping[str, object] | None,
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
|
|
@ -31,6 +32,7 @@ from litellm.router_utils.pre_call_checks.io_token_rate_limit_check import (
|
|||
io_token_reconcile_success,
|
||||
io_token_refund_failure,
|
||||
)
|
||||
from litellm.router_utils.routing_read_batch import ROUTER_USAGE_TARGET
|
||||
from litellm.types.router import RouterErrors
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
|
@ -137,6 +139,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
|
||||
return tpm_key, rpm_key
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def _get_current_tpm(self, tpm_key: str, tpm_limit: int) -> int | None:
|
||||
local_tpm: Final = self.dual_cache.get_cache(key=tpm_key, local_only=True)
|
||||
redis_cache: Final = self.dual_cache.redis_cache
|
||||
|
|
@ -147,6 +150,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
except RedisCircuitBreakerOpenError:
|
||||
return local_tpm
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def _async_get_current_tpm(self, tpm_key: str, tpm_limit: int, parent_otel_span: Span | None) -> int | None:
|
||||
local_tpm: Final = await self.dual_cache.async_get_cache(key=tpm_key, local_only=True)
|
||||
redis_cache: Final = self.dual_cache.redis_cache
|
||||
|
|
@ -157,6 +161,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
except RedisCircuitBreakerOpenError:
|
||||
return local_tpm
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def pre_call_check(self, deployment: dict) -> dict | None:
|
||||
"""
|
||||
Synchronous pre-call check for model rate limits.
|
||||
|
|
@ -236,6 +241,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
# Don't fail the request if rate limit check fails
|
||||
return deployment
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def async_pre_call_check(self, deployment: dict, parent_otel_span: Span | None = None) -> dict | None:
|
||||
"""
|
||||
Async pre-call check for model rate limits.
|
||||
|
|
@ -323,6 +329,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
# Don't fail the request if rate limit check fails
|
||||
return deployment
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
|
|
@ -394,6 +401,7 @@ class ModelRateLimitingCheck(CustomLogger):
|
|||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
)
|
||||
|
||||
@with_service_target(ROUTER_USAGE_TARGET)
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Sync version of tracking TPM usage after successful request.
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from pydantic import JsonValue, TypeAdapter
|
|||
from pydantic_core import to_jsonable_python
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm._internal_context import service_target
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.constants import PROMPT_CACHE_LOOKBACK_POSITIONS
|
||||
from litellm.litellm_core_utils.logging_utils import truncate_base64_in_messages
|
||||
|
|
@ -36,6 +37,7 @@ class PromptCachingCacheValue(TypedDict):
|
|||
|
||||
|
||||
PROMPT_CACHE_PIN_TTL_SECONDS: Final = 300
|
||||
_PROMPT_CACHE_PINS_TARGET: Final = "prompt_cache_pins"
|
||||
_TOOL_RUN_BLOCK_TYPES: Final = frozenset({"tool_use", "tool_result"})
|
||||
_PREFIX_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, JsonValue], ...])
|
||||
_TOOLS_ADAPTER: Final = TypeAdapter(tuple[JsonValue, ...])
|
||||
|
|
@ -291,11 +293,12 @@ class PromptCachingCache:
|
|||
if not positions:
|
||||
return
|
||||
|
||||
await self.cache.async_set_cache(
|
||||
positions[-1].cache_key,
|
||||
PromptCachingCacheValue(model_id=model_id),
|
||||
ttl=PROMPT_CACHE_PIN_TTL_SECONDS,
|
||||
)
|
||||
with service_target(_PROMPT_CACHE_PINS_TARGET):
|
||||
await self.cache.async_set_cache(
|
||||
positions[-1].cache_key,
|
||||
PromptCachingCacheValue(model_id=model_id),
|
||||
ttl=PROMPT_CACHE_PIN_TTL_SECONDS,
|
||||
)
|
||||
|
||||
async def async_get_model_id(
|
||||
self,
|
||||
|
|
@ -311,13 +314,9 @@ class PromptCachingCache:
|
|||
if not cache_keys:
|
||||
return None
|
||||
|
||||
return _first_pin(
|
||||
_PINS_ADAPTER.validate_python(
|
||||
await self.cache.async_batch_get_cache(
|
||||
keys=list(cache_keys),
|
||||
)
|
||||
)
|
||||
)
|
||||
with service_target(_PROMPT_CACHE_PINS_TARGET):
|
||||
pins: Final = await self.cache.async_batch_get_cache(keys=list(cache_keys))
|
||||
return _first_pin(_PINS_ADAPTER.validate_python(pins))
|
||||
|
||||
def get_model_id(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -17,11 +17,12 @@ from dataclasses import dataclass
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._internal_context import service_target
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_batch import BatchResult, active_request_redis_batches
|
||||
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
|
||||
from litellm.router_utils.cooldown_cache import CooldownCache
|
||||
from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET, CooldownCache
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span
|
||||
|
|
@ -29,9 +30,19 @@ if TYPE_CHECKING:
|
|||
from litellm.router import Router
|
||||
|
||||
|
||||
ROUTER_COOLDOWNS_USAGE_TARGET: Final = "router_cooldowns_usage"
|
||||
ROUTER_USAGE_TARGET: Final = "router_usage"
|
||||
_PREFETCH_SLOT: Final = "routing_read"
|
||||
|
||||
|
||||
def _routing_read_target(cooldown_keys: Sequence[str], usage_keys: Sequence[str]) -> str:
|
||||
if not usage_keys:
|
||||
return ROUTER_COOLDOWNS_TARGET
|
||||
if not cooldown_keys:
|
||||
return ROUTER_USAGE_TARGET
|
||||
return ROUTER_COOLDOWNS_USAGE_TARGET
|
||||
|
||||
|
||||
async def _backfill_prefetched_cache(
|
||||
cache: DualCache,
|
||||
due_keys: tuple[str, ...],
|
||||
|
|
@ -114,7 +125,8 @@ class RoutingPrefetch:
|
|||
)
|
||||
if not due:
|
||||
return
|
||||
result: Final = request.batch(redis_cache).mget(due)
|
||||
with service_target(_routing_read_target(cooldown_due, usage_due)):
|
||||
result: Final = request.batch(redis_cache).mget(due)
|
||||
prefetch: Final = RoutingPrefetch(
|
||||
keys=frozenset(keys), fetched=frozenset(due), result=result, reservations=reservations
|
||||
)
|
||||
|
|
@ -192,9 +204,10 @@ class RoutingReadBatch:
|
|||
(litellm_router_instance.cooldown_cache.cooldown_store, cooldown_keys),
|
||||
*(() if selector is None else ((selector.router_cache, list(usage_keys)),)),
|
||||
)
|
||||
results: Final = await self._read_prefetched(reads) or await DualCache.async_batch_get_cache_shared(
|
||||
reads, parent_otel_span=parent_otel_span
|
||||
)
|
||||
with service_target(_routing_read_target(cooldown_keys, usage_keys)):
|
||||
results: Final = await self._read_prefetched(reads) or await DualCache.async_batch_get_cache_shared(
|
||||
reads, parent_otel_span=parent_otel_span
|
||||
)
|
||||
cooldown_results: Final = results[0]
|
||||
if selector is not None:
|
||||
usage_values: Final = results[1]
|
||||
|
|
|
|||
|
|
@ -5,9 +5,12 @@ from typing import Final
|
|||
from pydantic import BaseModel
|
||||
|
||||
from litellm import print_verbose
|
||||
from litellm._internal_context import with_service_target
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.constants import DEFAULT_IN_MEMORY_TTL, DEFAULT_POLLING_INTERVAL
|
||||
|
||||
SCHEDULER_QUEUE_TARGET: Final = "scheduler_queue"
|
||||
|
||||
|
||||
class SchedulerCacheKeys(enum.Enum):
|
||||
queue = "scheduler:queue"
|
||||
|
|
@ -115,6 +118,7 @@ class Scheduler:
|
|||
"""Get the status of items in the queue"""
|
||||
return self.queue
|
||||
|
||||
@with_service_target(SCHEDULER_QUEUE_TARGET)
|
||||
async def get_queue(self, model_name: str) -> list:
|
||||
"""
|
||||
Return a queue for that specific model group
|
||||
|
|
@ -128,6 +132,7 @@ class Scheduler:
|
|||
return response
|
||||
return self.queue
|
||||
|
||||
@with_service_target(SCHEDULER_QUEUE_TARGET)
|
||||
async def save_queue(self, queue: list, model_name: str) -> None:
|
||||
"""
|
||||
Save the updated queue of the model group
|
||||
|
|
|
|||
|
|
@ -101,6 +101,7 @@ class ServiceLoggerPayload(BaseModel):
|
|||
duration: float = Field(description="How long did the request take?")
|
||||
call_type: str = Field(description="The call of the service, being made")
|
||||
caller: str | None = Field(None, description="The litellm call chain that made the service call, innermost first")
|
||||
target: str | None = Field(None, description="The key family the call served, e.g. llm_response or auth_objects")
|
||||
event_metadata: dict | None = Field(description="The metadata logged during service success/failure")
|
||||
|
||||
def to_json(self, **kwargs):
|
||||
|
|
|
|||
|
|
@ -91,6 +91,8 @@ ignored_function_names = [
|
|||
"_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py
|
||||
"_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name)
|
||||
"arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name)
|
||||
"_async_get_available_deployment", # Body of the `route {model}` phase wrapper, exercised through async_get_available_deployment in test_router.py
|
||||
"_async_get_available_deployment_for_pass_through", # Same, through async_get_available_deployment_for_pass_through in test_router.py
|
||||
"_embedding",
|
||||
"_aembedding",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -8,8 +8,10 @@ import pytest
|
|||
|
||||
import litellm
|
||||
import litellm.caching.redis_cache as redis_cache_module
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm._internal_context import current_service_target
|
||||
from litellm.caching.caching import Cache, response_cache_phase
|
||||
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.caching.redis_cache import RedisCache, _RedisTimeoutLogThrottle
|
||||
from litellm.types.caching import EMBEDDING_CACHE_FORMAT_VERSION, LiteLLMCacheType, SemanticCacheScope
|
||||
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
|
||||
|
|
@ -51,9 +53,7 @@ def test_cache_key_debug_log_does_not_include_prompt_material(caplog):
|
|||
assert re.fullmatch(r"[0-9a-f]{64}", cache_key)
|
||||
|
||||
created_cache_key_logs = [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if "Created cache key:" in record.getMessage()
|
||||
record.getMessage() for record in caplog.records if "Created cache key:" in record.getMessage()
|
||||
]
|
||||
assert created_cache_key_logs
|
||||
assert all(prompt_marker not in message for message in created_cache_key_logs)
|
||||
|
|
@ -86,13 +86,8 @@ def test_add_cache_timeout_only_joins_redis_throttle_for_redis_backends(backend,
|
|||
def _embedding_response(prompt_tokens, num_items):
|
||||
return EmbeddingResponse(
|
||||
model="amazon.titan-embed-image-v1",
|
||||
data=[
|
||||
Embedding(embedding=[0.0], index=i, object="embedding")
|
||||
for i in range(num_items)
|
||||
],
|
||||
usage=Usage(
|
||||
prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens
|
||||
),
|
||||
data=[Embedding(embedding=[0.0], index=i, object="embedding") for i in range(num_items)],
|
||||
usage=Usage(prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -144,9 +139,7 @@ def test_semantic_cache_key_excludes_prompt_so_paraphrases_share_a_bucket():
|
|||
)
|
||||
key_b = cache.get_cache_key(
|
||||
model="gpt-4o-mini",
|
||||
messages=[
|
||||
{"role": "user", "content": "Tell me the colour of the daytime sky."}
|
||||
],
|
||||
messages=[{"role": "user", "content": "Tell me the colour of the daytime sky."}],
|
||||
metadata=dict(tenant),
|
||||
)
|
||||
assert key_a == key_b
|
||||
|
|
@ -155,12 +148,8 @@ def test_semantic_cache_key_excludes_prompt_so_paraphrases_share_a_bucket():
|
|||
def test_semantic_cache_key_isolates_tenants():
|
||||
messages = [{"role": "user", "content": "What color is the sky?"}]
|
||||
cache = _semantic_cache()
|
||||
key_a = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, metadata={"user_api_key": "hash-A"}
|
||||
)
|
||||
key_b = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, metadata={"user_api_key": "hash-B"}
|
||||
)
|
||||
key_a = cache.get_cache_key(model="gpt-4o-mini", messages=messages, metadata={"user_api_key": "hash-A"})
|
||||
key_b = cache.get_cache_key(model="gpt-4o-mini", messages=messages, metadata={"user_api_key": "hash-B"})
|
||||
key_team = cache.get_cache_key(
|
||||
model="gpt-4o-mini",
|
||||
messages=messages,
|
||||
|
|
@ -244,24 +233,18 @@ def test_semantic_cache_key_still_separates_models_and_params():
|
|||
cache = _semantic_cache()
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
tenant = {"user_api_key": "hash-A"}
|
||||
assert cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, metadata=dict(tenant)
|
||||
) != cache.get_cache_key(model="gpt-4o", messages=messages, metadata=dict(tenant))
|
||||
assert cache.get_cache_key(model="gpt-4o-mini", messages=messages, metadata=dict(tenant)) != cache.get_cache_key(
|
||||
model="gpt-4o", messages=messages, metadata=dict(tenant)
|
||||
)
|
||||
assert cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, temperature=0, metadata=dict(tenant)
|
||||
) != cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=messages, temperature=1, metadata=dict(tenant)
|
||||
)
|
||||
) != cache.get_cache_key(model="gpt-4o-mini", messages=messages, temperature=1, metadata=dict(tenant))
|
||||
|
||||
|
||||
def test_exact_cache_key_still_includes_prompt():
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
key_a = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=[{"role": "user", "content": "a"}]
|
||||
)
|
||||
key_b = cache.get_cache_key(
|
||||
model="gpt-4o-mini", messages=[{"role": "user", "content": "b"}]
|
||||
)
|
||||
key_a = cache.get_cache_key(model="gpt-4o-mini", messages=[{"role": "user", "content": "a"}])
|
||||
key_b = cache.get_cache_key(model="gpt-4o-mini", messages=[{"role": "user", "content": "b"}])
|
||||
assert key_a != key_b
|
||||
|
||||
|
||||
|
|
@ -279,9 +262,7 @@ def test_exact_cache_key_includes_anthropic_messages_params(anthropic_param):
|
|||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
messages = [{"role": "user", "content": "which greek letter?"}]
|
||||
baseline = cache.get_cache_key(model="claude-sonnet-4-5", messages=messages)
|
||||
assert baseline != cache.get_cache_key(
|
||||
model="claude-sonnet-4-5", messages=messages, **anthropic_param
|
||||
)
|
||||
assert baseline != cache.get_cache_key(model="claude-sonnet-4-5", messages=messages, **anthropic_param)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -376,7 +357,9 @@ async def test_embedding_cache_serves_base64_string_embeddings_on_repeat(monkeyp
|
|||
self.provider_calls += 1
|
||||
return EmbeddingResponse(
|
||||
model=model,
|
||||
data=[Embedding(embedding="AACAPwAAAEA=", index=idx, object="embedding") for idx, _ in enumerate(input)],
|
||||
data=[
|
||||
Embedding(embedding="AACAPwAAAEA=", index=idx, object="embedding") for idx, _ in enumerate(input)
|
||||
],
|
||||
)
|
||||
|
||||
embedder = Base64Embedder()
|
||||
|
|
@ -403,3 +386,90 @@ def test_provider_specific_cache_key_ignores_litellm_owned_kwargs(monkeypatch: p
|
|||
assert cache.get_cache_key(**request, _litellm_control={"stream_chunk_size": 64}) == base_key
|
||||
assert cache.get_cache_key(**request, litellm_trace_id="trace-1") == base_key
|
||||
assert cache.get_cache_key(**{**request, "top_k": 6}) != base_key
|
||||
|
||||
|
||||
class PhaseRecordingCache(InMemoryCache):
|
||||
"""Records the target and the active span each read / write ran under, as a Redis span would."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.seen: list[tuple[str | None, str]] = []
|
||||
|
||||
def _record(self) -> None:
|
||||
from opentelemetry import trace
|
||||
|
||||
span = trace.get_current_span()
|
||||
self.seen.append((current_service_target(), getattr(span, "name", "")))
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
self._record()
|
||||
return super().get_cache(key, **kwargs)
|
||||
|
||||
def set_cache(self, key, value, **kwargs):
|
||||
self._record()
|
||||
super().set_cache(key, value, **kwargs)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def v2_span_exporter(monkeypatch):
|
||||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||
|
||||
from litellm.integrations.otel import OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.logger import OpenTelemetryV2
|
||||
from litellm.integrations.otel.plumbing import providers
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
config = OpenTelemetryV2Config(exporter="in_memory")
|
||||
exporter = InMemorySpanExporter()
|
||||
logger = OpenTelemetryV2(config=config, tracer_provider=providers.build_tracer_provider(config, exporter=exporter))
|
||||
monkeypatch.setattr(proxy_server, "open_telemetry_logger", logger)
|
||||
return exporter
|
||||
|
||||
|
||||
_REQUEST: Final = {"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "phase me"}]}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_facade_lookup_and_store_run_inside_the_response_cache_phases(v2_span_exporter):
|
||||
"""The native bridge calls ``Cache.async_get_cache`` / ``async_add_cache`` straight, never through
|
||||
``caching_handler``, so the ``cache.get llm_response`` / ``cache.set llm_response`` phase and the
|
||||
``llm_response`` target come from the facade: the store runs under them too, and a hit reads back."""
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
backend = PhaseRecordingCache()
|
||||
assert await cache.async_get_cache(dynamic_cache_object=backend, **_REQUEST) is None
|
||||
await cache.async_add_cache({"id": "resp-1"}, dynamic_cache_object=backend, **_REQUEST)
|
||||
assert await cache.async_get_cache(dynamic_cache_object=backend, **_REQUEST) == {"id": "resp-1"}
|
||||
assert backend.seen == [
|
||||
("llm_response", "cache.get llm_response"),
|
||||
("llm_response", "cache.set llm_response"),
|
||||
("llm_response", "cache.get llm_response"),
|
||||
]
|
||||
assert [s.name for s in v2_span_exporter.get_finished_spans()] == [
|
||||
"cache.get llm_response",
|
||||
"cache.set llm_response",
|
||||
"cache.get llm_response",
|
||||
]
|
||||
assert current_service_target() is None
|
||||
|
||||
|
||||
def test_sync_facade_lookup_and_store_run_inside_the_response_cache_phases(v2_span_exporter):
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
backend = PhaseRecordingCache()
|
||||
assert cache.get_cache(dynamic_cache_object=backend, **_REQUEST) is None
|
||||
cache.add_cache({"id": "resp-1"}, **_REQUEST)
|
||||
assert backend.seen == [("llm_response", "cache.get llm_response")]
|
||||
assert [s.name for s in v2_span_exporter.get_finished_spans()] == [
|
||||
"cache.get llm_response",
|
||||
"cache.set llm_response",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_lookup_already_inside_the_phase_does_not_open_a_second_one(v2_span_exporter):
|
||||
"""``caching_handler`` opens the phase around the facade call; the facade joins it."""
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
backend = PhaseRecordingCache()
|
||||
with response_cache_phase("get"):
|
||||
await cache.async_get_cache(dynamic_cache_object=backend, **_REQUEST)
|
||||
assert backend.seen == [("llm_response", "cache.get llm_response")]
|
||||
assert [s.name for s in v2_span_exporter.get_finished_spans()] == ["cache.get llm_response"]
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ import json
|
|||
import httpx
|
||||
import respx
|
||||
from fastapi.testclient import TestClient
|
||||
from litellm._internal_context import in_post_response_phase
|
||||
from litellm._internal_context import current_service_target, in_post_response_phase
|
||||
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
|
||||
|
||||
|
||||
|
|
@ -2268,3 +2268,48 @@ async def test_partial_embedding_cache_hit_sends_only_misses_and_keeps_input_ord
|
|||
|
||||
assert len(embedder.provider_inputs) == 2, embedder.provider_inputs
|
||||
assert [item["embedding"] for item in repeat.data] == [[float(len(text))] for text in mixed_input]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_cache_lookup_and_write_declare_the_llm_response_target(monkeypatch):
|
||||
"""Both the lookup and the write run under ``service_target("llm_response")`` so the
|
||||
datastore spans they issue read ``redis.get llm_response`` / ``redis.set llm_response``
|
||||
rather than by the cache method name."""
|
||||
seen: dict[str, str | None] = {}
|
||||
|
||||
class _TargetRecordingCache:
|
||||
supported_call_types = ["acompletion"]
|
||||
cache = None
|
||||
|
||||
def get_cache_key(self, **kwargs):
|
||||
return "k"
|
||||
|
||||
def _supports_async(self):
|
||||
return True
|
||||
|
||||
async def async_get_cache(self, **kwargs):
|
||||
seen["get"] = current_service_target()
|
||||
return None
|
||||
|
||||
async def async_add_cache(self, result, dynamic_cache_object=None, **kwargs):
|
||||
seen["set"] = current_service_target()
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
return None
|
||||
|
||||
handler = LLMCachingHandler(original_function=acompletion, request_kwargs={}, start_time=datetime.now())
|
||||
monkeypatch.setattr(litellm, "cache", _TargetRecordingCache())
|
||||
|
||||
await handler._async_get_cache(
|
||||
model="gpt-3.5-turbo",
|
||||
original_function=acompletion,
|
||||
logging_obj=MagicMock(),
|
||||
start_time=datetime.now(),
|
||||
call_type=CallTypes.acompletion.value,
|
||||
kwargs={"messages": [{"role": "user", "content": "hi"}]},
|
||||
)
|
||||
await handler.async_set_cache(result=litellm.ModelResponse(), original_function=acompletion, kwargs={})
|
||||
await asyncio.gather(*_PENDING_CACHE_WRITES)
|
||||
|
||||
assert seen == {"get": "llm_response", "set": "llm_response"}
|
||||
assert current_service_target() is None
|
||||
|
|
|
|||
|
|
@ -5,20 +5,26 @@ from __future__ import annotations
|
|||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Callable, Sequence
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from datetime import timedelta
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from redis.exceptions import NoScriptError
|
||||
|
||||
from litellm._internal_context import current_service_target, service_target
|
||||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.caching.redis_batch import (
|
||||
MIXED_PIPELINE_TARGET,
|
||||
RedisBatch,
|
||||
active_request_redis_batch,
|
||||
request_redis_batch_scope,
|
||||
)
|
||||
from litellm.caching.redis_cache import RedisCache, RedisCircuitBreaker
|
||||
from litellm.caching.redis_cache import (
|
||||
RedisCache,
|
||||
RedisCircuitBreaker,
|
||||
_get_call_stack_info, # pyright: ignore[reportPrivateUsage] # the chain the service hook reports
|
||||
)
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
|
||||
SCRIPT = "return redis.call('GET', KEYS[1])"
|
||||
|
|
@ -150,6 +156,30 @@ async def run_alone_script(keys: Sequence[str], args: Sequence[Any]) -> object:
|
|||
return ["alone", *keys, *args]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pipeline_flush_reports_its_name_as_the_call_type_and_the_op_count_as_metadata() -> None:
|
||||
"""The service event is ``request_redis_batch`` with ``op_count`` on the metadata, not
|
||||
``request_redis_batch[3]``: the span renders as ``redis.pipeline`` and the metrics label
|
||||
stays one value per batch name instead of one per batch size."""
|
||||
cache, _client = make()
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
async def record(**kwargs: Any) -> None:
|
||||
events.append(kwargs)
|
||||
|
||||
cache.service_logger_obj.async_service_success_hook = record # pyright: ignore[reportAttributeAccessIssue] # fake, records the hook call
|
||||
batch = RedisBatch(cache, name="request_redis_batch")
|
||||
got = batch.mget(["a:hit"])
|
||||
incr = batch.increment("cnt", 1)
|
||||
await got
|
||||
await incr
|
||||
await asyncio.gather(*(t for t in asyncio.all_tasks() if t is not asyncio.current_task()))
|
||||
|
||||
(event,) = events
|
||||
assert event["call_type"] == "request_redis_batch"
|
||||
assert event["event_metadata"] == {"op_count": 2}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_pipeline_carries_every_declared_operation_and_awaiting_one_flushes_all() -> None:
|
||||
cache, client = make(namespace="ns")
|
||||
|
|
@ -359,3 +389,122 @@ async def test_a_failed_mget_marks_nothing_as_missing() -> None:
|
|||
with pytest.raises(ConnectionError):
|
||||
await batch.mget(["b-miss"])
|
||||
assert batch.read_as_missing("b-miss") is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_operation_retried_alone_keeps_the_target_it_was_declared_under() -> None:
|
||||
"""The retry runs on the flush, outside the declaring caller's block, so the op carries
|
||||
the target it was declared under and the retried call is still named by its purpose."""
|
||||
seen: list[str | None] = []
|
||||
|
||||
async def record_target(keys: Sequence[str], args: Sequence[Any]) -> object:
|
||||
seen.append(current_service_target())
|
||||
return ["alone", *keys]
|
||||
|
||||
def reply_for(command: tuple[Any, ...]) -> Any:
|
||||
if command[0] == "EVALSHA":
|
||||
return NoScriptError("NOSCRIPT")
|
||||
return replies(command)
|
||||
|
||||
cache = FakeRedisCache(FakeClient(reply_for))
|
||||
batch = RedisBatch(cache)
|
||||
with service_target("spend_counters"):
|
||||
script = batch.script(SCRIPT, record_target, ["w"], [])
|
||||
assert current_service_target() is None
|
||||
assert await script == ["alone", "w"]
|
||||
assert seen == ["spend_counters"]
|
||||
assert current_service_target() is None
|
||||
|
||||
|
||||
class CallerRecordingClusterCache(FakeClusterCache):
|
||||
def __init__(self, client: FakeClient) -> None:
|
||||
super().__init__(client)
|
||||
self.callers: list[str] = []
|
||||
|
||||
async def async_batch_get_cache(self, key_list: Sequence[str], **kwargs: object) -> dict[str, Any]: # pyright: ignore[reportIncompatibleMethodOverride] # records what the service hook would report
|
||||
self.callers.append(_get_call_stack_info())
|
||||
return await super().async_batch_get_cache(key_list, **kwargs)
|
||||
|
||||
|
||||
def _prefetch_auth_objects(batch: RedisBatch) -> Awaitable[Sequence[Any]]:
|
||||
return batch.mget(["team", "user"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_cluster_op_names_the_code_that_declared_it_not_its_wrappers() -> None:
|
||||
"""On a cluster client every op runs alone, in a task driven by the flush, so above its
|
||||
wrappers there is only the event loop. Production reported ``_run_under_circuit_breaker <-
|
||||
wrapper``; the op carries the chain captured where it was declared and reports that."""
|
||||
cache = CallerRecordingClusterCache(FakeClient(replies))
|
||||
batch = RedisBatch(cache)
|
||||
with service_target("auth_objects"):
|
||||
pending = _prefetch_auth_objects(batch)
|
||||
assert await pending == {"team": None, "user": None}
|
||||
assert cache.callers == [
|
||||
"_prefetch_auth_objects <- test_a_cluster_op_names_the_code_that_declared_it_not_its_wrappers"
|
||||
]
|
||||
|
||||
|
||||
async def _flush_and_record_service_events(
|
||||
cache: FakeRedisCache, *results: Awaitable[object]
|
||||
) -> list[dict[str, object]]:
|
||||
events: list[dict[str, object]] = [] # mutable-ok: filled by the recording hooks
|
||||
|
||||
async def record(**kwargs: object) -> None:
|
||||
events.append({**kwargs, "target": current_service_target()})
|
||||
|
||||
cache.service_logger_obj.async_service_success_hook = record # pyright: ignore[reportAttributeAccessIssue] # fake, records the hook call
|
||||
cache.service_logger_obj.async_service_failure_hook = record # pyright: ignore[reportAttributeAccessIssue] # fake, records the hook call
|
||||
await asyncio.gather(*results, return_exceptions=True)
|
||||
await asyncio.gather(*(t for t in asyncio.all_tasks() if t is not asyncio.current_task()))
|
||||
return events
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pipeline_of_one_key_family_is_targeted_by_that_family() -> None:
|
||||
"""Every op in the flush was declared under ``auth_objects``, so the span is
|
||||
``redis.pipeline auth_objects`` and carries only the op count."""
|
||||
cache, _client = make()
|
||||
batch = RedisBatch(cache, name="request_redis_batch")
|
||||
with service_target("auth_objects"):
|
||||
first = batch.mget(["a:hit"])
|
||||
second = batch.mget(["b:hit"])
|
||||
|
||||
(event,) = await _flush_and_record_service_events(cache, first, second)
|
||||
assert (event["target"], event["event_metadata"]) == ("auth_objects", {"op_count": 2})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pipeline_of_several_key_families_is_mixed_and_lists_the_families_sorted() -> None:
|
||||
"""Owners of different families sharing one round trip render as ``redis.pipeline mixed``
|
||||
with the sorted family list beside the op count, never as a bare ``redis.pipeline``."""
|
||||
cache, _client = make()
|
||||
batch = RedisBatch(cache, name="request_redis_batch")
|
||||
with service_target("spend_counters"):
|
||||
incr = batch.increment("cnt", 1)
|
||||
with service_target("auth_objects"):
|
||||
auth = batch.mget(["a:hit"])
|
||||
with service_target("router_cooldowns"):
|
||||
cooldown = batch.mget(["c:hit"])
|
||||
|
||||
(event,) = await _flush_and_record_service_events(cache, incr, auth, cooldown)
|
||||
assert event["target"] == MIXED_PIPELINE_TARGET
|
||||
assert event["event_metadata"] == {"op_count": 3, "families": "auth_objects,router_cooldowns,spend_counters"}
|
||||
assert current_service_target() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_pipeline_reports_the_same_family_target_as_a_successful_one() -> None:
|
||||
"""The failure event names the pipeline the same way, so the error span lines up with the
|
||||
success spans of the same flush shape in a trace search."""
|
||||
cache, _client = make(fail=ConnectionError("redis down"))
|
||||
batch = RedisBatch(cache, name="post_call_redis_batch")
|
||||
with service_target("spend_counters"):
|
||||
incr = batch.increment("cnt", 1)
|
||||
with service_target("auth_objects"):
|
||||
auth = batch.mget(["a:hit"])
|
||||
|
||||
(event,) = await _flush_and_record_service_events(cache, incr, auth)
|
||||
assert isinstance(event["error"], ConnectionError)
|
||||
assert (event["call_type"], event["target"]) == ("post_call_redis_batch", MIXED_PIPELINE_TARGET)
|
||||
assert event["event_metadata"] == {"op_count": 2, "families": "auth_objects,spend_counters"}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import time
|
||||
import types
|
||||
from collections.abc import Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Final
|
||||
|
|
@ -59,9 +60,7 @@ def test_check_and_fix_namespace_prefixes_keys_sharing_the_namespace_prefix(
|
|||
|
||||
@pytest.mark.parametrize("namespace", [None, "litellm"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_applies_namespace(
|
||||
namespace, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_delete_cache_applies_namespace(namespace, monkeypatch, redis_no_ping):
|
||||
"""async_delete_cache must prefix keys with the namespace, matching every
|
||||
other cache operation. Without this, Redis NOPERM errors occur when an
|
||||
ACL restricts DEL to the litellm:* pattern."""
|
||||
|
|
@ -69,9 +68,7 @@ async def test_async_delete_cache_applies_namespace(
|
|||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_delete_cache(key="3997c4abcdef")
|
||||
|
||||
expected_key = "litellm:3997c4abcdef" if namespace else "3997c4abcdef"
|
||||
|
|
@ -134,9 +131,7 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
|
|||
]
|
||||
|
||||
# Test the helper method
|
||||
result = await redis_cache.handle_lpop_count_for_older_redis_versions(
|
||||
pipe=mock_pipeline, key="test_key", count=2
|
||||
)
|
||||
result = await redis_cache.handle_lpop_count_for_older_redis_versions(pipe=mock_pipeline, key="test_key", count=2)
|
||||
|
||||
# Verify results
|
||||
assert result == [b"value1", b"value2"]
|
||||
|
|
@ -145,18 +140,14 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_pipeline_empty_list_returns_empty(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_rpush_pipeline_empty_list_returns_empty(monkeypatch, redis_no_ping):
|
||||
"""Empty rpush_list should return empty list without touching Redis"""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache()
|
||||
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
result = await redis_cache.async_rpush_pipeline(rpush_list=[])
|
||||
|
||||
assert result == []
|
||||
|
|
@ -171,9 +162,7 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
|
|||
|
||||
mock_redis_instance = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
result = await redis_cache.async_lpop_pipeline(lpop_list=[])
|
||||
|
||||
assert result == []
|
||||
|
|
@ -198,9 +187,7 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
|
|||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_namespaces_keys(
|
||||
namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_namespaces_keys(namespace, raw_keys, expected_keys, monkeypatch, redis_no_ping):
|
||||
"""The callable returned by async_register_script (used by the rate limiter
|
||||
Lua scripts, pod-lock release, and budget limiters) must namespace every key
|
||||
it is invoked with. The hash tag is preserved so cluster slotting is intact."""
|
||||
|
|
@ -211,16 +198,12 @@ async def test_async_register_script_namespaces_keys(
|
|||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.register_script = MagicMock(return_value=registered_script)
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
script = redis_cache.async_register_script("return 1")
|
||||
result = await script(keys=raw_keys, args=[60])
|
||||
|
||||
assert result == "ok"
|
||||
registered_script.assert_awaited_once_with(
|
||||
keys=tuple(expected_keys), args=[60], client=None
|
||||
)
|
||||
registered_script.assert_awaited_once_with(keys=tuple(expected_keys), args=[60], client=None)
|
||||
|
||||
|
||||
# LIT-3298: rate limits tripped at ~40M instead of 80M. async_register_script
|
||||
|
|
@ -258,12 +241,8 @@ def test_async_register_script_binds_per_event_loop(namespace, monkeypatch):
|
|||
loop_a = asyncio.new_event_loop()
|
||||
loop_b = asyncio.new_event_loop()
|
||||
try:
|
||||
result_a = loop_a.run_until_complete(
|
||||
script(keys=["{k:v}:tokens"], args=[60])
|
||||
)
|
||||
result_b = loop_b.run_until_complete(
|
||||
script(keys=["{k:v}:tokens"], args=[60])
|
||||
)
|
||||
result_a = loop_a.run_until_complete(script(keys=["{k:v}:tokens"], args=[60]))
|
||||
result_b = loop_b.run_until_complete(script(keys=["{k:v}:tokens"], args=[60]))
|
||||
finally:
|
||||
loop_a.close()
|
||||
loop_b.close()
|
||||
|
|
@ -276,9 +255,7 @@ def test_async_register_script_binds_per_event_loop(namespace, monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_not_shared_across_namespaces(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_not_shared_across_namespaces(monkeypatch, redis_no_ping):
|
||||
"""Two caches with different namespaces registering the SAME script must
|
||||
each run against their own client and key prefix. A content-only executor
|
||||
cache would let the second cache reuse the first's executor and namespace."""
|
||||
|
|
@ -294,9 +271,10 @@ async def test_async_register_script_not_shared_across_namespaces(
|
|||
client_b.register_script = MagicMock(return_value=reg_b)
|
||||
|
||||
same_script = "return redis.call('GET', KEYS[1])"
|
||||
with patch.object(
|
||||
cache_a, "init_async_client", return_value=client_a
|
||||
), patch.object(cache_b, "init_async_client", return_value=client_b):
|
||||
with (
|
||||
patch.object(cache_a, "init_async_client", return_value=client_a),
|
||||
patch.object(cache_b, "init_async_client", return_value=client_b),
|
||||
):
|
||||
script_a = cache_a.async_register_script(same_script)
|
||||
script_b = cache_b.async_register_script(same_script)
|
||||
result_a = await script_a(keys=["k"], args=[])
|
||||
|
|
@ -308,9 +286,7 @@ async def test_async_register_script_not_shared_across_namespaces(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_cluster_path_uses_evalsha(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_cluster_path_uses_evalsha(monkeypatch, redis_no_ping):
|
||||
"""Redis Cluster exposes script_load/evalsha rather than register_script.
|
||||
The script is loaded once and invoked via evalsha with namespaced keys."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
|
|
@ -320,23 +296,17 @@ async def test_async_register_script_cluster_path_uses_evalsha(
|
|||
cluster_client.script_load = MagicMock(return_value="sha123")
|
||||
cluster_client.evalsha = AsyncMock(return_value="cluster-ok")
|
||||
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=cluster_client
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=cluster_client):
|
||||
script = redis_cache.async_register_script("return 'cluster'")
|
||||
result = await script(keys=["{k:v}:tokens"], args=[5, 60])
|
||||
|
||||
assert result == "cluster-ok"
|
||||
cluster_client.script_load.assert_called_once_with("return 'cluster'")
|
||||
cluster_client.evalsha.assert_awaited_once_with(
|
||||
"sha123", 1, "ns:{k:v}:tokens", 5, 60
|
||||
)
|
||||
cluster_client.evalsha.assert_awaited_once_with("sha123", 1, "ns:{k:v}:tokens", 5, 60)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_register_script_raises_for_unsupported_client(
|
||||
monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_register_script_raises_for_unsupported_client(monkeypatch, redis_no_ping):
|
||||
"""A client exposing neither register_script nor script_load fails loudly
|
||||
rather than silently returning a no-op callable."""
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
|
|
@ -351,46 +321,34 @@ async def test_async_register_script_raises_for_unsupported_client(
|
|||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_delete_cache_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_delete_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_delete_cache("k")
|
||||
mock_redis_instance.delete.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_cache_keys_namespaces_keys(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_delete_cache_keys_namespaces_keys(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.delete_cache_keys(["k"])
|
||||
mock_redis_instance.delete.assert_awaited_once_with(expected)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_ttl_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_get_ttl_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.ttl = AsyncMock(return_value=42)
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
ttl = await redis_cache.async_get_ttl("k")
|
||||
assert ttl == 42
|
||||
mock_redis_instance.ttl.assert_awaited_once_with(expected)
|
||||
|
|
@ -398,41 +356,31 @@ async def test_async_get_ttl_namespaces_key(
|
|||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_lpop_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_lpop_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.lpop = AsyncMock(return_value=b"value")
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_lpop(key="k")
|
||||
mock_redis_instance.lpop.assert_awaited_once_with(expected, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_rpush_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_rpush_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_redis_instance = AsyncMock()
|
||||
mock_redis_instance.rpush = AsyncMock(return_value=1)
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_rpush("k", ["v"])
|
||||
mock_redis_instance.rpush.assert_awaited_once_with(expected, "v")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected_match", [(None, "k*"), ("ns", "ns:k*")])
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_scan_iter_namespaces_pattern(
|
||||
namespace, expected_match, monkeypatch, redis_no_ping
|
||||
):
|
||||
async def test_async_scan_iter_namespaces_pattern(namespace, expected_match, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
|
||||
|
|
@ -449,17 +397,13 @@ async def test_async_scan_iter_namespaces_pattern(
|
|||
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.scan_iter = scan_iter
|
||||
with patch.object(
|
||||
redis_cache, "init_async_client", return_value=mock_redis_instance
|
||||
):
|
||||
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
|
||||
await redis_cache.async_scan_iter(pattern="k")
|
||||
assert captured["match"] == expected_match
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace, expected", [(None, "k"), ("ns", "ns:k")])
|
||||
def test_increment_cache_namespaces_key(
|
||||
namespace, expected, monkeypatch, redis_no_ping
|
||||
):
|
||||
def test_increment_cache_namespaces_key(namespace, expected, monkeypatch, redis_no_ping):
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
redis_cache = RedisCache(namespace=namespace)
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -1534,7 +1478,7 @@ class _ListPipeline:
|
|||
self.rows.extend(op[2:])
|
||||
results.append(len(self.rows))
|
||||
else:
|
||||
start, end = int(op[2]), int(op[3])
|
||||
start = int(op[2])
|
||||
del self.rows[: max(len(self.rows) + start, 0) if start < 0 else start]
|
||||
results.append(True)
|
||||
return results
|
||||
|
|
@ -1556,3 +1500,114 @@ async def test_async_rpush_and_trim_runs_push_and_trim_in_one_transaction(monkey
|
|||
assert pushed_len == 4
|
||||
assert rows == ["b", "c", "d"]
|
||||
assert pipe.queued == [("rpush", "ns:buf", "c", "d"), ("ltrim", "ns:buf", "-3", "-1")]
|
||||
|
||||
|
||||
def test_call_stack_info_skips_generic_cache_facade_frames():
|
||||
"""A read through ``DualCache.async_get_cache`` -> ``RedisCache.async_get_cache`` used to
|
||||
report ``async_get_cache <- async_get_cache``; the chain names the code that wanted the
|
||||
read, skipping the facade verbs and the batch retry wrappers in between."""
|
||||
from litellm.caching.redis_cache import _get_call_stack_info
|
||||
|
||||
def probe(): # the RedisCache method that sets call_type
|
||||
return _get_call_stack_info()
|
||||
|
||||
def async_get_cache(): # a facade's generic verb
|
||||
return probe()
|
||||
|
||||
def run_alone(): # the batch retry wrapper
|
||||
return async_get_cache()
|
||||
|
||||
def _retrieve_from_cache():
|
||||
return run_alone()
|
||||
|
||||
def _async_get_cache():
|
||||
return _retrieve_from_cache()
|
||||
|
||||
assert _async_get_cache() == "_retrieve_from_cache <- _async_get_cache"
|
||||
|
||||
|
||||
def test_call_stack_info_stops_at_the_event_loop():
|
||||
"""Event-loop frames are not callers, so a read issued straight from a task names the
|
||||
task's coroutine alone rather than padding the chain with asyncio internals."""
|
||||
from litellm.caching.redis_cache import _get_call_stack_info
|
||||
|
||||
def probe():
|
||||
return _get_call_stack_info()
|
||||
|
||||
async def _lookup():
|
||||
return probe()
|
||||
|
||||
assert asyncio.run(_lookup()) == "_lookup"
|
||||
|
||||
|
||||
def test_call_stack_info_reports_the_threaded_caller_when_only_wrappers_are_found():
|
||||
"""A batch op retried on the flush runs in a task of its own, so above its wrappers there
|
||||
is only the event loop; the chain is the one its declaring code threaded through
|
||||
``service_caller``, never the wrapper names (``run_alone <- _settle_alone`` says nothing)."""
|
||||
from litellm._internal_context import service_caller
|
||||
from litellm.caching.redis_cache import _get_call_stack_info
|
||||
|
||||
def probe():
|
||||
return _get_call_stack_info()
|
||||
|
||||
def run_alone():
|
||||
return probe()
|
||||
|
||||
async def _settle_alone():
|
||||
return run_alone()
|
||||
|
||||
async def flush():
|
||||
with service_caller("prefetch_auth_objects <- user_api_key_auth"):
|
||||
task = asyncio.create_task(_settle_alone())
|
||||
return await task
|
||||
|
||||
assert asyncio.run(flush()) == "prefetch_auth_objects <- user_api_key_auth"
|
||||
|
||||
|
||||
def test_call_stack_info_is_unknown_when_only_wrappers_are_found_and_nothing_was_threaded():
|
||||
import threading
|
||||
|
||||
from litellm.caching.redis_cache import _get_call_stack_info
|
||||
|
||||
def probe():
|
||||
return _get_call_stack_info()
|
||||
|
||||
def run_alone():
|
||||
return probe()
|
||||
|
||||
def _settle_alone():
|
||||
return run_alone()
|
||||
|
||||
seen: list[str] = []
|
||||
worker = threading.Thread(target=lambda: seen.append(_settle_alone()))
|
||||
worker.start()
|
||||
worker.join()
|
||||
assert seen == ["unknown"]
|
||||
|
||||
|
||||
def _native_probe():
|
||||
from litellm.caching.redis_cache import _get_call_stack_info
|
||||
|
||||
return _get_call_stack_info()
|
||||
|
||||
|
||||
def _settle():
|
||||
return _native_probe()
|
||||
|
||||
|
||||
def drive():
|
||||
return _settle()
|
||||
|
||||
|
||||
def test_call_stack_info_skips_native_lifecycle_frames():
|
||||
"""The Rust execution awaits the response-cache coroutine from ``lifecycle._settle`` inside
|
||||
``drive``; those frames forward every native suspension, so the chain names the code that
|
||||
started the native call instead of ``_settle <- drive``."""
|
||||
lifecycle_globals = {"__name__": "litellm.rust_bridge.lifecycle", "_native_probe": _native_probe}
|
||||
native_settle = types.FunctionType(_settle.__code__, lifecycle_globals, "_settle")
|
||||
native_drive = types.FunctionType(drive.__code__, {**lifecycle_globals, "_settle": native_settle}, "drive")
|
||||
|
||||
def anthropic_messages():
|
||||
return native_drive()
|
||||
|
||||
assert anthropic_messages() == "anthropic_messages <- test_call_stack_info_skips_native_lifecycle_frames"
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue