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:
devin-ai-integration[bot] 2026-10-03 09:20:27 -07:00 • committed by GitHub
parent 53d2ab6b05
commit 564d236985
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
114 changed files with 3070 additions and 987 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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, ...]]":

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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