From 564d236985352589e53b1e8565acc7d74b787178 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 09:20:27 -0700 Subject: [PATCH] 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 Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../enterprise_hooks/blocked_user_list.py | 19 +- .../send_emails/base_email.py | 8 +- .../proxy/hooks/managed_files.py | 8 + litellm/_internal_context.py | 83 +- litellm/_service_logger.py | 3 + litellm/caching/affinity_cache.py | 2 + litellm/caching/caching.py | 235 +++--- litellm/caching/caching_handler.py | 55 +- litellm/caching/redis_batch.py | 65 +- litellm/caching/redis_cache.py | 98 ++- .../SlackAlerting/hanging_request_check.py | 4 + .../SlackAlerting/slack_alerting.py | 31 +- litellm/integrations/custom_guardrail.py | 5 + litellm/integrations/otel/README.md | 100 ++- litellm/integrations/otel/logger.py | 39 +- litellm/integrations/otel/mappers/genai.py | 12 +- litellm/integrations/otel/model/metadata.py | 38 +- litellm/integrations/otel/model/payloads.py | 31 +- litellm/integrations/otel/model/semconv.py | 6 + litellm/integrations/otel/model/spans.py | 54 +- litellm/integrations/otel/plumbing/context.py | 58 +- litellm/integrations/prometheus.py | 3 + litellm/litellm_core_utils/litellm_logging.py | 27 +- .../base_managed_resource.py | 6 + .../mcp_server/auth/user_api_key_auth_mcp.py | 4 + .../mcp_server/gateway_dcr_flow.py | 5 + .../mcp_server/mcp_server_manager.py | 6 + .../mcp_server/oauth2_token_cache.py | 4 + .../dual_cache_token_backend.py | 5 + .../proxy/_experimental/mcp_server/utils.py | 3 + .../proxy/agent_endpoints/identity_store.py | 4 + .../anthropic_endpoints/gateway_endpoints.py | 4 + litellm/proxy/auth/auth_checks.py | 37 + litellm/proxy/auth/auth_object_prefetch.py | 28 +- litellm/proxy/auth/handle_jwt.py | 7 + litellm/proxy/auth/login_throttle.py | 6 + litellm/proxy/auth/resolvers/store.py | 3 + litellm/proxy/auth/user_api_key_auth.py | 17 +- .../auth_cache_invalidation_pubsub.py | 4 + .../proxy/common_utils/reset_budget_job.py | 10 +- .../proxy/common_utils/user_api_key_cache.py | 2 + litellm/proxy/db/db_spend_update_writer.py | 17 +- .../db_transaction_queue/pod_lock_manager.py | 5 + .../redis_update_buffer.py | 11 + litellm/proxy/db/gateway_request_tracking.py | 5 + litellm/proxy/db/spend_counter_reseed.py | 10 +- .../guardrails/guardrail_hooks/lasso/lasso.py | 3 + .../shared_health_check_manager.py | 9 + litellm/proxy/hooks/batch_enqueued_tokens.py | 5 + litellm/proxy/hooks/batch_rate_limiter.py | 3 + litellm/proxy/hooks/batch_redis_get.py | 20 +- litellm/proxy/hooks/dynamic_rate_limiter.py | 6 + .../proxy/hooks/dynamic_rate_limiter_v3.py | 4 + .../hooks/max_budget_per_session_limiter.py | 4 + litellm/proxy/hooks/max_iterations_limiter.py | 2 + .../proxy/hooks/model_max_budget_limiter.py | 9 + .../proxy/hooks/parallel_request_limiter.py | 8 + .../hooks/parallel_request_limiter_v3.py | 16 + .../proxy/hooks/prompt_cache_prediction.py | 2 + litellm/proxy/hooks/sensitive_data_routing.py | 4 + .../access_group_endpoints.py | 4 + .../internal_user_endpoints.py | 3 + .../key_management_endpoints.py | 9 +- .../mcp_management_endpoints.py | 4 + .../management_endpoints/sso/saml_sso.py | 5 + .../management_endpoints/sso_helper_utils.py | 5 + litellm/proxy/management_endpoints/ui_sso.py | 14 +- litellm/proxy/management_helpers/utils.py | 6 + litellm/proxy/proxy_server.py | 108 ++- .../proxy/response_polling/polling_handler.py | 7 + .../spend_tracking/budget_reservation.py | 5 + .../spend_tracking/ptu_flat_cost_rollup.py | 3 + .../spend_tracking/spend_capture_rate.py | 3 + .../spend_tracking/spend_counter_batch.py | 11 +- .../proxy_setting_endpoints.py | 6 +- litellm/proxy/utils.py | 26 +- litellm/router.py | 55 +- .../router_strategy/base_routing_strategy.py | 3 + litellm/router_strategy/budget_limiter.py | 7 + .../complexity_router/complexity_router.py | 4 +- litellm/router_strategy/least_busy.py | 8 + litellm/router_strategy/lowest_cost.py | 4 + litellm/router_strategy/lowest_latency.py | 6 + litellm/router_strategy/lowest_tpm_rpm.py | 4 + litellm/router_strategy/lowest_tpm_rpm_v2.py | 7 + litellm/router_utils/cooldown_cache.py | 24 +- litellm/router_utils/cooldown_handlers.py | 15 +- litellm/router_utils/health_state_cache.py | 7 + .../deployment_affinity_check.py | 9 +- .../io_token_rate_limit_check.py | 10 + .../pre_call_checks/model_rate_limit_check.py | 8 + litellm/router_utils/prompt_caching_cache.py | 23 +- litellm/router_utils/routing_read_batch.py | 23 +- litellm/scheduler.py | 5 + litellm/types/services.py | 1 + .../router_code_coverage.py | 2 + tests/unit/caching/test_caching.py | 142 +++- tests/unit/caching/test_caching_handler.py | 47 +- tests/unit/caching/test_redis_batch.py | 153 +++- tests/unit/caching/test_redis_cache.py | 237 ++++-- .../test_request_redis_batch_post_call.py | 27 + .../test_request_redis_batch_pre_call.py | 52 +- .../SlackAlerting/test_slack_alerting.py | 30 + .../otel/test_otel_v2_components.py | 8 +- .../otel/test_otel_v2_destinations.py | 108 ++- .../integrations/otel/test_otel_v2_logger.py | 782 ++++++++++-------- .../hooks/test_sensitive_data_routing.py | 177 ++-- ...test_proxy_server_endpoints_and_startup.py | 45 + .../test_config_param_cache.py | 29 + .../unit/router_utils/test_cooldown_cache.py | 47 +- .../router_utils/test_cooldown_handlers.py | 34 + tests/unit/test_internal_context.py | 261 ++++++ tests/unit/test_router/test_router.py | 109 +++ tests/unit/test_service_logger.py | 43 +- 114 files changed, 3070 insertions(+), 987 deletions(-) create mode 100644 tests/unit/test_internal_context.py diff --git a/enterprise/enterprise_hooks/blocked_user_list.py b/enterprise/enterprise_hooks/blocked_user_list.py index a032ea7662d..dfaf91ea081 100644 --- a/enterprise/enterprise_hooks/blocked_user_list.py +++ b/enterprise/enterprise_hooks/blocked_user_list.py @@ -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, diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index 6e33d9f1bf3..a29f0a1b43a 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -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) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 8ca8857445c..e0a94612646 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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: diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py index 389add8ed0f..df549701cf4 100644 --- a/litellm/_internal_context.py +++ b/litellm/_internal_context.py @@ -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.""" diff --git a/litellm/_service_logger.py b/litellm/_service_logger.py index 1a5f46e9261..24625382d0c 100644 --- a/litellm/_service_logger.py +++ b/litellm/_service_logger.py @@ -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, ) diff --git a/litellm/caching/affinity_cache.py b/litellm/caching/affinity_cache.py index 3387d4953a0..2d27574b098 100644 --- a/litellm/caching/affinity_cache.py +++ b/litellm/caching/affinity_cache.py @@ -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 = """ diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 9e04ca79822..12aaa051041 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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) diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index ee022822872..f04ff8b6e78 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -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 diff --git a/litellm/caching/redis_batch.py b/litellm/caching/redis_batch.py index b3596aab6a1..aa458c9c926 100644 --- a/litellm/caching/redis_batch.py +++ b/litellm/caching/redis_batch.py @@ -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 diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 29e390b1d9a..2ee4bff6112 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -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( - {"", "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( + { + "", + "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}") diff --git a/litellm/integrations/SlackAlerting/hanging_request_check.py b/litellm/integrations/SlackAlerting/hanging_request_check.py index 4d7cbfe8fd1..6f986144d4c 100644 --- a/litellm/integrations/SlackAlerting/hanging_request_check.py +++ b/litellm/integrations/SlackAlerting/hanging_request_check.py @@ -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 diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 6ff048c484d..4c2b722ef90 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -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 diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 02ac53a541b..67e2173ecd0 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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, diff --git a/litellm/integrations/otel/README.md b/litellm/integrations/otel/README.md index d9047b675ce..cdd1b9ef95b 100644 --- a/litellm/integrations/otel/README.md +++ b/litellm/integrations/otel/README.md @@ -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 diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index e1983a44451..740580dfcf3 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -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: diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index e37da8908e4..7947dbaae29 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -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 diff --git a/litellm/integrations/otel/model/metadata.py b/litellm/integrations/otel/model/metadata.py index 7cb64debfe0..809fc794461 100644 --- a/litellm/integrations/otel/model/metadata.py +++ b/litellm/integrations/otel/model/metadata.py @@ -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. diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index e06cc1d0407..efd7f6c7dd9 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -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 diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 19b319009e8..2d05754d37b 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -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`` diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 35fc50a2a83..6ddc824235b 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -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]: diff --git a/litellm/integrations/otel/plumbing/context.py b/litellm/integrations/otel/plumbing/context.py index 19356939046..d0c419abe87 100644 --- a/litellm/integrations/otel/plumbing/context.py +++ b/litellm/integrations/otel/plumbing/context.py @@ -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, ...]]": diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 0a14cd7cf18..6468dc41ea5 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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, diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index d8e70d3b125..26c02bb0243 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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: diff --git a/litellm/llms/base_llm/managed_resources/base_managed_resource.py b/litellm/llms/base_llm/managed_resources/base_managed_resource.py index 4fbc0ce51b0..81b3411ed56 100644 --- a/litellm/llms/base_llm/managed_resources/base_managed_resource.py +++ b/litellm/llms/base_llm/managed_resources/base_managed_resource.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 9c7778e2b77..f575265a5e7 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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. diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index e66504af47a..1136410fd18 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -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. diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 816ae9aea59..a776470e2ab 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 3742d7b4ccc..080665fda8a 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py index 361a9632dc1..69ff7cda2b7 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/dual_cache_token_backend.py @@ -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)) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 7c9d75457b5..464e3043458 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py index 3c8163a8838..c3e2064cbd5 100644 --- a/litellm/proxy/agent_endpoints/identity_store.py +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -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 diff --git a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py index 6fec411313d..e84bb73b05b 100644 --- a/litellm/proxy/anthropic_endpoints/gateway_endpoints.py +++ b/litellm/proxy/anthropic_endpoints/gateway_endpoints.py @@ -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 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index dbd6f28a183..918428bcf0e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index f8b0a3838f0..d4a91d1f194 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -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( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 4448d860217..cfa9d10b74c 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -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, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index f2398219a7b..0536e293f13 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -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: diff --git a/litellm/proxy/auth/resolvers/store.py b/litellm/proxy/auth/resolvers/store.py index 832baf03432..13e24387551 100644 --- a/litellm/proxy/auth/resolvers/store.py +++ b/litellm/proxy/auth/resolvers/store.py @@ -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() diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e82f3eed7cc..950aab723d4 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index 3e09ad7157f..519e604783e 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -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) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index c4f081e3dec..72cac1224f6 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -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: diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index c99665986dd..8c830816207 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -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}") diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 26a21069c83..d9a60742c9d 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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 diff --git a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py index bc67617e444..627b30e83f7 100644 --- a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py +++ b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py @@ -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. diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 534ba30a6d0..3d24b0a0612 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -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, diff --git a/litellm/proxy/db/gateway_request_tracking.py b/litellm/proxy/db/gateway_request_tracking.py index aae15f81b06..2dbcd1ccdf0 100644 --- a/litellm/proxy/db/gateway_request_tracking.py +++ b/litellm/proxy/db/gateway_request_tracking.py @@ -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 diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index f8e102d2682..721999b1868 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -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", diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index 985812ca980..e05db5003fa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -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. diff --git a/litellm/proxy/health_check_utils/shared_health_check_manager.py b/litellm/proxy/health_check_utils/shared_health_check_manager.py index 79d54df97ae..764d570fbea 100644 --- a/litellm/proxy/health_check_utils/shared_health_check_manager.py +++ b/litellm/proxy/health_check_utils/shared_health_check_manager.py @@ -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. diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 1a410593854..7b4fa6fe625 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -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, diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index fc2f97ca57e..894123b256a 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -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, diff --git a/litellm/proxy/hooks/batch_redis_get.py b/litellm/proxy/hooks/batch_redis_get.py index 13e2bdbc304..3d4ef67bb91 100644 --- a/litellm/proxy/hooks/batch_redis_get.py +++ b/litellm/proxy/hooks/batch_redis_get.py @@ -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) diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index f4eac6ae5ae..8c41eb8d2d3 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -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): diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 0339cf4dfea..d600e249754 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -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. diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index e07b96e5773..2ba42c43dcc 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -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: diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index 93697afa3c6..efcafc1b6b0 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -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, diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index d2db5145fce..d019271d404 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -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 diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index e3485ebf25d..b4ce010dd27 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -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. diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b33cea5742d..2bfaf57f0fc 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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, diff --git a/litellm/proxy/hooks/prompt_cache_prediction.py b/litellm/proxy/hooks/prompt_cache_prediction.py index 65c456c5666..e724d973d95 100644 --- a/litellm/proxy/hooks/prompt_cache_prediction.py +++ b/litellm/proxy/hooks/prompt_cache_prediction.py @@ -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: diff --git a/litellm/proxy/hooks/sensitive_data_routing.py b/litellm/proxy/hooks/sensitive_data_routing.py index bc89dec7a11..1773fc2d50a 100644 --- a/litellm/proxy/hooks/sensitive_data_routing.py +++ b/litellm/proxy/hooks/sensitive_data_routing.py @@ -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, diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index fb53c06928a..97311a0ef8a 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -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, diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 04b8ec56ae2..d0b3a08bc77 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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. diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 88627cdb5aa..323b9e434a8 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 0d358302fa2..be4e46f247e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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: diff --git a/litellm/proxy/management_endpoints/sso/saml_sso.py b/litellm/proxy/management_endpoints/sso/saml_sso.py index 12e1f1a03f3..292ba3c5988 100644 --- a/litellm/proxy/management_endpoints/sso/saml_sso.py +++ b/litellm/proxy/management_endpoints/sso/saml_sso.py @@ -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, diff --git a/litellm/proxy/management_endpoints/sso_helper_utils.py b/litellm/proxy/management_endpoints/sso_helper_utils.py index 11f4184437b..2cf27a254ae 100644 --- a/litellm/proxy/management_endpoints/sso_helper_utils.py +++ b/litellm/proxy/management_endpoints/sso_helper_utils.py @@ -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""" diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 0d7094f66dc..01807fefd78 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -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. diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index 0698556f3bc..b8af3950859 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9af080f1b50..0e88ff9c13f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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(): diff --git a/litellm/proxy/response_polling/polling_handler.py b/litellm/proxy/response_polling/polling_handler.py index fe7fa79a3d9..4ae227162d2 100644 --- a/litellm/proxy/response_polling/polling_handler.py +++ b/litellm/proxy/response_polling/polling_handler.py @@ -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: diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 5eed1a894bc..2dd1c9419f1 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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, diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index 99e5c35ae14..b2bf6f46b3a 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -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. diff --git a/litellm/proxy/spend_tracking/spend_capture_rate.py b/litellm/proxy/spend_tracking/spend_capture_rate.py index bd804afa4ff..664f100a7b2 100644 --- a/litellm/proxy/spend_tracking/spend_capture_rate.py +++ b/litellm/proxy/spend_tracking/spend_capture_rate.py @@ -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( diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index ae24331c236..7e0821e091e 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -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 diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 610d47990c3..a112b22b4e4 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -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), diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5828ed94984..eed2c257e87 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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" diff --git a/litellm/router.py b/litellm/router.py index 24554e61516..cbc7e655810 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/litellm/router_strategy/base_routing_strategy.py b/litellm/router_strategy/base_routing_strategy.py index 79d457f2836..51400e6da24 100644 --- a/litellm/router_strategy/base_routing_strategy.py +++ b/litellm/router_strategy/base_routing_strategy.py @@ -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: diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index a792654e2a9..631b0c3df3d 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -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: diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 3610991a20d..fc0865654c7 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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, diff --git a/litellm/router_strategy/least_busy.py b/litellm/router_strategy/least_busy.py index 9ab670e4b95..6f3e0936641 100644 --- a/litellm/router_strategy/least_busy.py +++ b/litellm/router_strategy/least_busy.py @@ -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: diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 22c321c65fb..d567b6acccc 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -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, diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 66c8227195d..622919e3443 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -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, diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index d4abf1f8f70..2d373e0c266 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -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, diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 25564a80e0a..9839c9be469 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -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, diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index 187215d3d16..f780bb3364a 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -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 diff --git a/litellm/router_utils/cooldown_handlers.py b/litellm/router_utils/cooldown_handlers.py index 408ddbab34b..fcafdfb7402 100644 --- a/litellm/router_utils/cooldown_handlers.py +++ b/litellm/router_utils/cooldown_handlers.py @@ -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( diff --git a/litellm/router_utils/health_state_cache.py b/litellm/router_utils/health_state_cache.py index c8ca7105392..4fb9476dae0 100644 --- a/litellm/router_utils/health_state_cache.py +++ b/litellm/router_utils/health_state_cache.py @@ -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: diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index edba4c27647..432fe11dc2e 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -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, diff --git a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py index fbd3e18e357..0d701411c94 100644 --- a/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/io_token_rate_limit_check.py @@ -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, diff --git a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py index 79ea6dc36ec..3f911bf0825 100644 --- a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py @@ -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. diff --git a/litellm/router_utils/prompt_caching_cache.py b/litellm/router_utils/prompt_caching_cache.py index 23005a97c59..e6a8dc88086 100644 --- a/litellm/router_utils/prompt_caching_cache.py +++ b/litellm/router_utils/prompt_caching_cache.py @@ -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, diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py index 752b2857de4..7f6267ec9b9 100644 --- a/litellm/router_utils/routing_read_batch.py +++ b/litellm/router_utils/routing_read_batch.py @@ -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] diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 028b5d085e2..e19e386f527 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -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 diff --git a/litellm/types/services.py b/litellm/types/services.py index b8c4265b6be..00fa9f044cc 100644 --- a/litellm/types/services.py +++ b/litellm/types/services.py @@ -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): diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 8d7c1e140d2..f55f415b76c 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -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", ] diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 0e0f2b7eac6..0a7ac3ecad1 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -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"] diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index 6cf8e901cd7..1599668839a 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -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 diff --git a/tests/unit/caching/test_redis_batch.py b/tests/unit/caching/test_redis_batch.py index 93206efc80f..cd270035416 100644 --- a/tests/unit/caching/test_redis_batch.py +++ b/tests/unit/caching/test_redis_batch.py @@ -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"} diff --git a/tests/unit/caching/test_redis_cache.py b/tests/unit/caching/test_redis_cache.py index 5f83be7c7bc..5db11a67564 100644 --- a/tests/unit/caching/test_redis_cache.py +++ b/tests/unit/caching/test_redis_cache.py @@ -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" diff --git a/tests/unit/caching/test_request_redis_batch_post_call.py b/tests/unit/caching/test_request_redis_batch_post_call.py index fdd328a9a57..8c8c9df5926 100644 --- a/tests/unit/caching/test_request_redis_batch_post_call.py +++ b/tests/unit/caching/test_request_redis_batch_post_call.py @@ -15,6 +15,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest import litellm +from litellm._internal_context import current_service_target from litellm.caching.caching import Cache from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache @@ -27,6 +28,7 @@ from litellm.caching.redis_batch import ( ) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.proxy.auth.auth_object_prefetch import AUTH_OBJECTS_TARGET from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PARALLEL_RELEASE_SCRIPT, TOKEN_INCREMENT_SCRIPT, @@ -541,6 +543,31 @@ async def test_the_update_cache_read_armed_before_accounting_rides_the_pipeline_ assert active_request_redis_batches() is None +@pytest.mark.asyncio +async def test_the_armed_update_cache_read_is_declared_under_the_auth_objects_family(): + """The user, team and tag rows the accounting reads are auth objects, so the pipeline that carries + the armed read renders ``redis.pipeline auth_objects``, not a bare ``redis.pipeline``.""" + from litellm.proxy.proxy_server import _read_update_cache_values, arm_update_cache_read + + client = FakeClient(_ok_replies) + redis_cache = PostCallFakeRedisCache(client) + cache = DualCache() + cache.attach_redis_cache(redis_cache) + pipeline_targets: list[str | None] = [] # mutable-ok: filled by the recording hook + + async def record(**kwargs: object) -> None: + pipeline_targets.append(current_service_target()) + + redis_cache.service_logger_obj.async_service_success_hook = record # pyright: ignore[reportAttributeAccessIssue] # fake, records the hook call + + with request_redis_batch_scope(): + await arm_update_cache_read(["user-1", "team_id:t1"], cache=cache) + await _read_update_cache_values(["user-1", "team_id:t1"], None, cache=cache) + await asyncio.gather(*(t for t in asyncio.all_tasks() if t is not asyncio.current_task())) + + assert pipeline_targets == [AUTH_OBJECTS_TARGET] + + @pytest.mark.asyncio async def test_an_update_cache_read_armed_for_other_keys_is_ignored_and_the_read_happens_as_before(): from litellm.proxy.proxy_server import _read_update_cache_values, arm_update_cache_read diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py index d4388110131..cee7bfd8c65 100644 --- a/tests/unit/caching/test_request_redis_batch_pre_call.py +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -12,8 +12,9 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from litellm import Router import litellm.caching.dual_cache as dual_cache_module +from litellm import Router +from litellm._internal_context import current_service_target from litellm.caching.dual_cache import DualCache from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable @@ -27,8 +28,13 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache -from litellm.router_utils.cooldown_cache import CooldownCache -from litellm.router_utils.routing_read_batch import RoutingPrefetch +from litellm.router_utils.cooldown_cache import ROUTER_COOLDOWNS_TARGET, CooldownCache +from litellm.router_utils.routing_read_batch import ( + ROUTER_COOLDOWNS_USAGE_TARGET, + ROUTER_USAGE_TARGET, + RoutingPrefetch, + _routing_read_target, # pyright: ignore[reportPrivateUsage] # the family rule under test +) from .test_redis_batch import FakeClient, FakeRedisCache, replies @@ -422,6 +428,30 @@ async def test_a_failed_prefetch_falls_back_to_the_shared_read(): assert len(fallback_cooldown_mgets) == 1 +@pytest.mark.asyncio +async def test_the_armed_routing_read_is_declared_under_the_router_cooldowns_family(): + """The prefetch is declared before routing runs under a target of its own, so the pipeline that + carries it renders ``redis.pipeline router_cooldowns`` instead of a bare ``redis.pipeline``.""" + client = FakeClient(_lua_ok_replies) + redis_cache = FakeRedisCache(client) + router = _router(redis_cache, routing_strategy="simple-shuffle") + pipeline_targets: list[str | None] = [] # mutable-ok: filled by the recording hook + + async def record(**kwargs: object) -> None: + pipeline_targets.append(current_service_target()) + + redis_cache.service_logger_obj.async_service_success_hook = record # pyright: ignore[reportAttributeAccessIssue] # fake, records the hook call + + with request_redis_batch_scope(): + router.arm_routing_read_prefetch(_MODEL_GROUP, {}) + await router.async_get_available_deployment( + model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={} + ) + await asyncio.gather(*(t for t in asyncio.all_tasks() if t is not asyncio.current_task())) + + assert pipeline_targets == [ROUTER_COOLDOWNS_TARGET] + + @pytest.mark.asyncio async def test_an_abandoned_prefetch_still_backfills_the_cooldown_it_read(monkeypatch): clock: Final = 1_000_000.0 @@ -1037,3 +1067,19 @@ async def test_identity_prefetch_is_one_mget_after_which_hits_and_misses_alike_c assert await cache.async_get_cache("end_user_id:eu-miss") is None assert len(client.pipelines) == 1 and redis_cache.alone == [] assert cache.in_memory_cache.get_cache("end_user_id:eu-miss") is None + + +@pytest.mark.parametrize( + ("cooldown_keys", "usage_keys", "expected"), + [ + (("cooldown:a",), (), ROUTER_COOLDOWNS_TARGET), + ((), ("usage:a",), ROUTER_USAGE_TARGET), + (("cooldown:a",), ("usage:a",), ROUTER_COOLDOWNS_USAGE_TARGET), + ], +) +def test_the_routing_read_family_follows_the_keys_that_are_actually_due( + cooldown_keys: tuple[str, ...], usage_keys: tuple[str, ...], expected: str +): + """A routing MGET is ``router_cooldowns`` when only cooldown keys go out, ``router_usage`` when the + cooldowns were already in memory and only usage counters go out, and the combined family otherwise.""" + assert _routing_read_target(cooldown_keys, usage_keys) == expected diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py index 0c2b95fd448..47e55c8476e 100644 --- a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py @@ -12,6 +12,7 @@ from pydantic import TypeAdapter from typing_extensions import ReadOnly, TypedDict import litellm +from litellm._internal_context import current_service_target from litellm.caching.caching import DualCache from litellm.integrations.SlackAlerting.budget_alert_types import get_budget_alert_type from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting @@ -579,3 +580,32 @@ async def test_update_values_repeated_alerting_reload_keeps_single_periodic_flus await t except asyncio.CancelledError: pass + + +@pytest.mark.asyncio +async def test_daily_report_schedule_cache_calls_declare_their_key_family(): + """The report_sent read and write run inside ``service_target("daily_report_schedule")`` + so the background spans read ``redis.get daily_report_schedule`` rather than a bare + ``redis.get`` with no owner.""" + slack_alerting: Final = await _slack_alerting_with_due_daily_report() + cache: Final = slack_alerting.internal_usage_cache + seen: list[tuple[str, str | None]] = [] + real_get, real_set = cache.async_get_cache, cache.async_set_cache + + async def _get(*args, **kwargs): + seen.append(("get", current_service_target())) + return await real_get(*args, **kwargs) + + async def _set(*args, **kwargs): + seen.append(("set", current_service_target())) + return await real_set(*args, **kwargs) + + with ( + patch.object(cache, "async_get_cache", side_effect=_get), + patch.object(cache, "async_set_cache", side_effect=_set), + ): + result: Final = await slack_alerting._run_scheduler_helper(llm_router=MagicMock(), pod_lock_manager=None) + + assert result is True + assert seen == [("get", "daily_report_schedule"), ("set", "daily_report_schedule")] + assert current_service_target() is None diff --git a/tests/unit/integrations/otel/test_otel_v2_components.py b/tests/unit/integrations/otel/test_otel_v2_components.py index fb7be0dda14..bc230a8bcb2 100644 --- a/tests/unit/integrations/otel/test_otel_v2_components.py +++ b/tests/unit/integrations/otel/test_otel_v2_components.py @@ -145,18 +145,21 @@ def test_service_span_data_from_payload(): service = _Service() call_type = "async_set_cache" caller = "async_set_cache <- async_add_cache" + target = "llm_response" error = None data = ServiceSpanData.from_payload(_Payload()) assert data.service_name == "redis" assert data.call_type == "async_set_cache" assert data.caller == "async_set_cache <- async_add_cache" + assert data.target == "llm_response" assert data.error is None class _FailPayload: service = _Service() call_type = "async_set_cache" caller = None + target = None error = "boom" failed = ServiceSpanData.from_payload(_FailPayload()) @@ -1429,7 +1432,7 @@ def test_sanitize_event_metadata_drops_objects_dumps_and_secrets(): clean = sanitize_event_metadata( { "table_name": "combined_view", # safe primitive -> kept - "count": 3, # primitive -> kept (stringified) + "count": 3, # primitive -> kept, still an int "function_kwargs": {"prisma_client": object()}, # denylisted key "function_args": (1, 2), # denylisted key "user_api_key_auth": "blob", # 'auth' substring -> dropped @@ -1440,7 +1443,8 @@ def test_sanitize_event_metadata_drops_objects_dumps_and_secrets(): "nested": {"x": 1}, # non-primitive value -> dropped } ) - assert clean == {"table_name": "combined_view", "count": "3"} + assert clean == {"table_name": "combined_view", "count": 3} + assert isinstance(clean["count"], int) def test_sanitize_event_metadata_caps_value_length_and_handles_none(): diff --git a/tests/unit/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py index cb986e61229..3fdc47f130c 100644 --- a/tests/unit/integrations/otel/test_otel_v2_destinations.py +++ b/tests/unit/integrations/otel/test_otel_v2_destinations.py @@ -1,9 +1,11 @@ """Key/team OTLP destinations override the operator's exporters for that backend.""" +import asyncio import contextvars import time from base64 import b64encode from collections.abc import Mapping +from datetime import datetime, timezone from functools import reduce from types import MappingProxyType @@ -1116,7 +1118,9 @@ class TestProviderWiring: assert self._fan_out_of(preset)._excluded_db_systems == frozenset({"redis"}) - @pytest.mark.parametrize("otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"]) + @pytest.mark.parametrize( + "otel", [None, True, "on", "", []], ids=["null", "true", "on", "empty_string", "empty_list"] + ) def test_a_non_mapping_otel_block_falls_back_to_the_published_logger_config(self, monkeypatch, otel): monkeypatch.setattr(litellm, "callback_settings", {"otel": otel}, raising=False) preset = OpenTelemetryV2( @@ -2062,6 +2066,28 @@ def credential_less_proxy(monkeypatch) -> None: langfuse_preset() +def _closed_chat_call_kwargs() -> dict[str, object]: + """The callback kwargs of one completed chat call, as both an operator and a destination logger see them.""" + payload = { + "call_type": "acompletion", + "custom_llm_provider": "openai", + "model": "gpt-4o", + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "stream": False, + "response": {"id": "resp_1", "model": "gpt-4o", "choices": [{"finish_reason": "stop"}]}, + "metadata": {"team_id": "t1", "user_api_key_hash": "hsh"}, + "status": "success", + "litellm_call_id": "call_dup_1", + } + return { + "standard_logging_object": payload, + "litellm_params": {"metadata": {}}, + "api_call_start_time": datetime(2026, 5, 26, 12, 0, 0, tzinfo=timezone.utc), + } + + class TestPresetDegradation: def test_a_credential_less_langfuse_exports_nowhere_instead_of_to_the_console(self, monkeypatch, capfd): """``_normalize`` folds a console exporter in for an empty list, which would @@ -2270,28 +2296,78 @@ class TestPresetDegradation: assert logger is None - def test_a_credentialed_logger_beside_another_v2_logger_keeps_every_exporter(self, monkeypatch): - """Only a degraded preset gives the collector up; an operator who configured - both the backend and the collector still exports to both, as on base.""" + @pytest.mark.parametrize("anchored", [True, False]) + def test_a_logger_built_beside_another_v2_logger_keeps_only_its_backends_exporter(self, monkeypatch, anchored): + """Operator credentials for the backend do not make the collector safe to copy: the + registered logger already exports every call there, so a copy of ``chat`` riding the + preset's base exporters lands in the operator's sink a second time. A key's ``logging`` + entry reaches this builder as a plain dynamic callback too, with no destination anchored.""" from litellm.litellm_core_utils.litellm_logging import _maybe_construct_otel_v2 - monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-lf-1") - monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-lf-1") - monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com") + for name in ("ARIZE_SPACE_KEY", "ARIZE_ENDPOINT", "ARIZE_HTTP_ENDPOINT", "ARIZE_PROJECT_NAME"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("ARIZE_SPACE_ID", "space-operator") + monkeypatch.setenv("ARIZE_API_KEY", "ak-operator") monkeypatch.setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://collector.local:4318") monkeypatch.setenv("LITELLM_OTEL_V2", "true") - collector_logger = build_otel_v2_logger(OpenTelemetryV2Config(exporter="in_memory")) + operator_exporter = InMemorySpanExporter() + operator_cfg = OpenTelemetryV2Config(exporter="in_memory") + operator = build_otel_v2_logger( + operator_cfg, tracer_provider=otel_providers.build_tracer_provider(operator_cfg, exporter=operator_exporter) + ) + + def run(): + if anchored: + set_request_destinations( + ( + OtelDestination( + endpoint="https://otlp.arize.com/v1", headers={"space_id": "t"}, callback_name="arize" + ), + ) + ) + return _maybe_construct_otel_v2("arize", [operator]) is_otel_v2_enabled.cache_clear() - logger = in_fresh_context(_maybe_construct_otel_v2, "langfuse_otel", [collector_logger]) + tenant = in_fresh_context(run) is_otel_v2_enabled.cache_clear() - assert logger is not None - assert [spec.endpoint for spec in logger.config.exporters] == [ - "http://collector.local:4318", - "https://cloud.langfuse.com/api/public/otel", - ] - assert all(spec.headers for spec in logger.config.exporters if spec.requires_headers) + assert tenant is not None + assert [spec.owner for spec in tenant.config.exporters] == [ExporterOwner.ARIZE_AX] + assert "http://collector.local:4318" not in {spec.endpoint for spec in tenant.config.exporters} + + tenant_exporter = InMemorySpanExporter() + tenant_twin = build_otel_v2_logger( + tenant.config, + callback_name="arize", + tracer_provider=otel_providers.build_tracer_provider(tenant.config, exporter=tenant_exporter), + ) + kwargs = _closed_chat_call_kwargs() + operator.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs) + asyncio.run(operator.async_log_success_event(kwargs, None, None, None)) + asyncio.run(tenant_twin.async_log_success_event(kwargs, None, None, None)) + assert [span.name for span in operator_exporter.get_finished_spans()] == ["chat gpt-4o"] + assert [span.name for span in tenant_exporter.get_finished_spans()] == ["chat gpt-4o"] + + def test_a_preset_that_owns_no_exporter_keeps_the_collector_it_was_built_on(self, monkeypatch): + """Langtrace is a mapper over the operator's own OTLP collector and contributes no exporter + of its own, so filtering to owned exporters would register it with nowhere to deliver.""" + from litellm.litellm_core_utils.litellm_logging import _maybe_construct_otel_v2 + + monkeypatch.setenv("OTEL_EXPORTER_OTLP_ENDPOINT", "http://collector.local:4318") + monkeypatch.setenv("LITELLM_OTEL_V2", "true") + operator_cfg = OpenTelemetryV2Config(exporter="in_memory") + operator = build_otel_v2_logger( + operator_cfg, + tracer_provider=otel_providers.build_tracer_provider(operator_cfg, exporter=InMemorySpanExporter()), + ) + + is_otel_v2_enabled.cache_clear() + langtrace = in_fresh_context(lambda: _maybe_construct_otel_v2("langtrace", [operator])) + is_otel_v2_enabled.cache_clear() + + assert langtrace is not None + assert "langtrace" in langtrace.config.mapper_names + assert "http://collector.local:4318" in {spec.endpoint for spec in langtrace.config.exporters} class TestContextIsolation: @@ -2856,8 +2932,8 @@ class TestEvictionSafety: import threading from litellm.integrations.otel.plumbing.providers import ( - _DrainPool, _MAX_CACHED_DESTINATION_PROCESSORS, + _DrainPool, ) class GatedDrain(_DrainPool): diff --git a/tests/unit/integrations/otel/test_otel_v2_logger.py b/tests/unit/integrations/otel/test_otel_v2_logger.py index 62bf75bd083..b7bf678d1fe 100644 --- a/tests/unit/integrations/otel/test_otel_v2_logger.py +++ b/tests/unit/integrations/otel/test_otel_v2_logger.py @@ -23,8 +23,15 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( # noqa: E4 from opentelemetry.trace import SpanKind # noqa: E402 from opentelemetry.trace.status import StatusCode # noqa: E402 -from litellm._internal_context import in_post_response_phase, post_response_phase # noqa: E402 -from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY # noqa: E402 +from litellm._internal_context import ( # noqa: E402 + in_post_response_phase, + post_response_phase, + service_target, +) +from litellm.constants import ( # noqa: E402 + INTERNAL_CALL_ORIGIN_METADATA_KEY, + SESSION_ID_GENERATED_METADATA_KEY, +) from litellm.integrations.otel import ( # noqa: E402 GenAI, LiteLLM, @@ -45,6 +52,7 @@ from litellm.integrations.otel.plumbing.context import ( # noqa: E402 set_mcp_message_transport_span, set_request_root_span, ) +from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN # noqa: E402 # --------------------------------------------------------------------------- # # Fixtures @@ -127,11 +135,7 @@ def _emit_llm(logger, kwargs=None, *, ambient=None, fail=False): if kwargs is None: kwargs = _kwargs() payload = kwargs.get("standard_logging_object") or {} - with ( - trace.use_span(ambient, end_on_exit=False) - if ambient is not None - else contextlib.nullcontext() - ): + with trace.use_span(ambient, end_on_exit=False) if ambient is not None else contextlib.nullcontext(): logger.log_pre_api_call(model=payload.get("model"), messages=[], kwargs=kwargs) hook = logger.async_log_failure_event if fail else logger.async_log_success_event asyncio.run(hook(kwargs, None, None, None)) @@ -409,9 +413,7 @@ def test_sync_log_event_is_noop(): def test_missing_standard_logging_object_is_noop(): """No carrier (``pre_call`` never ran) → the callback emits nothing.""" logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event({"litellm_params": {}}, None, None, None) - ) + asyncio.run(logger.async_log_success_event({"litellm_params": {}}, None, None, None)) assert exporter.get_finished_spans() == () @@ -426,9 +428,7 @@ def test_no_span_when_pre_call_never_ran(): error_information={"error_class": "ProxyException", "error_code": "401"}, ) # No log_pre_api_call: the call never started. - asyncio.run( - logger.async_log_failure_event(_kwargs(payload=payload), None, None, None) - ) + asyncio.run(logger.async_log_failure_event(_kwargs(payload=payload), None, None, None)) assert exporter.get_finished_spans() == () # no phantom LLM span @@ -598,11 +598,7 @@ def test_mcp_tool_call_stateless_omits_session_id(): logger, exporter = _logger() payload = _mcp_payload() del payload["metadata"]["mcp_tool_call_metadata"]["mcp_session_id"] - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": payload}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": payload}, None, None, None)) (span,) = exporter.get_finished_spans() assert "mcp.session.id" not in span.attributes assert span.attributes["mcp.method.name"] == "tools/call" @@ -636,11 +632,7 @@ def test_mcp_tool_call_failure_marks_error(): status="failure", error_information={"error_class": "MCPError", "error_message": "upstream 500"}, ) - asyncio.run( - logger.async_log_failure_event( - {"standard_logging_object": payload}, None, None, None - ) - ) + asyncio.run(logger.async_log_failure_event({"standard_logging_object": payload}, None, None, None)) (span,) = exporter.get_finished_spans() assert span.name == "tools/call get_weather" assert span.status.status_code is StatusCode.ERROR @@ -667,14 +659,8 @@ def test_mcp_tool_call_metadata_read_from_nested_metadata_not_top_level(): # Move the real metadata to the top level only, mirroring the old buggy read # location. ``call_type`` still classifies this as an MCP call, so the span is # emitted, but none of its fields are reachable from the wrong nesting level. - payload["mcp_tool_call_metadata"] = payload["metadata"].pop( - "mcp_tool_call_metadata" - ) - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": payload}, None, None, None - ) - ) + payload["mcp_tool_call_metadata"] = payload["metadata"].pop("mcp_tool_call_metadata") + asyncio.run(logger.async_log_success_event({"standard_logging_object": payload}, None, None, None)) (span,) = exporter.get_finished_spans() assert span.name == "tools/call" assert "mcp.session.id" not in span.attributes @@ -741,11 +727,7 @@ def test_mcp_tool_call_names_its_rpc_system_and_upstream(): dependency ``:0``, which is worse than leaving the span unclassified. """ logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": _mcp_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": _mcp_payload()}, None, None, None)) (span,) = exporter.get_finished_spans() assert span.attributes["rpc.system"] == "jsonrpc" assert span.attributes["server.address"] == "weather.example.com" @@ -774,9 +756,7 @@ def test_mcp_tool_call_omits_rpc_system_without_a_complete_upstream(resource): del payload["metadata"]["mcp_tool_call_metadata"]["mcp_server_resource"] else: payload["metadata"]["mcp_tool_call_metadata"]["mcp_server_resource"] = resource - asyncio.run( - logger.async_log_success_event({"standard_logging_object": payload}, None, None, None) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": payload}, None, None, None)) (span,) = exporter.get_finished_spans() assert "rpc.system" not in span.attributes assert "server.port" not in span.attributes @@ -792,20 +772,14 @@ def test_mcp_list_tools_omits_rpc_system_without_an_upstream(): dependency node in every consumer that aggregates on it. """ logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": _mcp_list_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": _mcp_list_payload()}, None, None, None)) (span,) = exporter.get_finished_spans() assert "rpc.system" not in span.attributes assert "server.address" not in span.attributes @pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES) -def test_mcp_span_nests_under_transport_without_propagated_context( - make_payload, span_name -): +def test_mcp_span_nests_under_transport_without_propagated_context(make_payload, span_name): """Almost no MCP client implements SEP-414, so ``params._meta`` normally carries no trace context. Rooting the span there split one tool call into two traces joined only by a link, which is how it surfaced in APM: the ``POST`` transaction @@ -813,15 +787,9 @@ def test_mcp_span_nests_under_transport_without_propagated_context( honor the span nests under the transport span instead, and records no link since the transport is now the real parent.""" logger, exporter = _logger() - transport = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + transport = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(transport) - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None)) transport.end() span = next(s for s in exporter.get_finished_spans() if s.name == span_name) assert span.parent is not None @@ -831,9 +799,7 @@ def test_mcp_span_nests_under_transport_without_propagated_context( @pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES) -def test_mcp_span_nests_under_this_messages_transport_not_the_session_opener( - make_payload, span_name -): +def test_mcp_span_nests_under_this_messages_transport_not_the_session_opener(make_payload, span_name): """A *stateful* streamable-HTTP session runs every message on the single task spawned by that session's ``initialize`` POST, so the ``_request_root_span`` ContextVar the ASGI request task writes is frozen at ``initialize`` inside the @@ -843,19 +809,13 @@ def test_mcp_span_nests_under_this_messages_transport_not_the_session_opener( current message's transport on the request task and publishes it, so the span parents to the POST that actually carried this message.""" logger, exporter = _logger() - session_opener = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) - this_message = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + session_opener = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + this_message = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) async def session_task(): token = set_mcp_message_transport_span(this_message) try: - await logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) + await logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None) finally: reset_mcp_message_transport_span(token) @@ -877,26 +837,18 @@ def test_mcp_span_nests_under_this_messages_transport_not_the_session_opener( @pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES) -def test_mcp_span_roots_without_transport_or_propagated_context( - make_payload, span_name -): +def test_mcp_span_roots_without_transport_or_propagated_context(make_payload, span_name): """With neither a remote parent nor a transport span there is nothing to nest under, so the span legitimately starts its own root trace with no links.""" logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None)) span = next(s for s in exporter.get_finished_spans() if s.name == span_name) assert span.parent is None assert span.links == () @pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES) -def test_mcp_span_links_propagated_meta_trace_context_and_nests_under_transport( - make_payload, span_name -): +def test_mcp_span_links_propagated_meta_trace_context_and_nests_under_transport(make_payload, span_name): """When the client propagates W3C trace context in the request's ``params._meta`` (SEP-414), the MCP span still nests under the gateway's own transport span — one renderable trace — and records the client's context as a @@ -904,19 +856,11 @@ def test_mcp_span_links_propagated_meta_trace_context_and_nests_under_transport( trace whose root span never reaches the gateway's tracing backend, leaving the span unreachable from the trace view.""" logger, exporter = _logger() - transport = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + transport = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(transport) - token = set_mcp_message_trace_carrier( - {"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"} - ) + token = set_mcp_message_trace_carrier({"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"}) try: - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None)) finally: reset_mcp_message_trace_carrier(token) transport.end() @@ -924,29 +868,19 @@ def test_mcp_span_links_propagated_meta_trace_context_and_nests_under_transport( assert span.parent is not None assert span.parent.span_id == transport.get_span_context().span_id assert span.context.trace_id == transport.get_span_context().trace_id - assert [link.context.trace_id for link in span.links] == [ - 0x11111111111111111111111111111111 - ] + assert [link.context.trace_id for link in span.links] == [0x11111111111111111111111111111111] assert [link.context.span_id for link in span.links] == [0x2222222222222222] @pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES) -def test_mcp_span_without_transport_roots_and_links_propagated_context( - make_payload, span_name -): +def test_mcp_span_without_transport_roots_and_links_propagated_context(make_payload, span_name): """With no transport span at all there is nothing of the gateway's to anchor to, so the span starts its own root trace — and the client context stays a span link there too, so the event keeps one shape everywhere.""" logger, exporter = _logger() - token = set_mcp_message_trace_carrier( - {"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"} - ) + token = set_mcp_message_trace_carrier({"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"}) try: - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None)) finally: reset_mcp_message_trace_carrier(token) span = next(s for s in exporter.get_finished_spans() if s.name == span_name) @@ -960,19 +894,11 @@ def test_mcp_span_links_unsampled_client_traceparent(): remote context, so the link is recorded; the span's own recording follows the transport's sampling decision, never the client's flag.""" logger, exporter = _logger() - transport = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + transport = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(transport) - token = set_mcp_message_trace_carrier( - {"traceparent": "00-11111111111111111111111111111111-2222222222222222-00"} - ) + token = set_mcp_message_trace_carrier({"traceparent": "00-11111111111111111111111111111111-2222222222222222-00"}) try: - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": _mcp_list_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": _mcp_list_payload()}, None, None, None)) finally: reset_mcp_message_trace_carrier(token) transport.end() @@ -992,9 +918,7 @@ def test_mcp_span_ignores_client_supplied_baggage(make_payload, span_name): extracts trace context only, so the spoofed keys never reach the span while the legitimate traceparent parenting still works.""" logger, exporter = _logger() - transport = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + transport = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(transport) token = set_mcp_message_trace_carrier( { @@ -1003,11 +927,7 @@ def test_mcp_span_ignores_client_supplied_baggage(make_payload, span_name): } ) try: - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None)) finally: reset_mcp_message_trace_carrier(token) transport.end() @@ -1029,11 +949,7 @@ def test_mcp_span_carries_authenticated_identity(make_payload, span_name): span — parented to an empty remote context — would carry no team/key attribute at all, so it couldn't be attributed or filtered by team in the traces backend.""" logger, exporter = _logger() - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": make_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": make_payload()}, None, None, None)) span = next(s for s in exporter.get_finished_spans() if s.name == span_name) assert span.attributes[LiteLLM.TEAM_ID] == "t1" @@ -1044,17 +960,11 @@ def test_mcp_span_malformed_traceparent_nests_under_transport(): falls back to nesting under the transport span rather than starting a disconnected root trace.""" logger, exporter = _logger() - transport = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + transport = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(transport) token = set_mcp_message_trace_carrier({"traceparent": "not-a-valid-traceparent"}) try: - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": _mcp_list_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": _mcp_list_payload()}, None, None, None)) finally: reset_mcp_message_trace_carrier(token) transport.end() @@ -1069,23 +979,15 @@ def test_mcp_span_with_propagated_context_nests_under_this_messages_transport(): carrying this message, not the stale session anchor — otherwise the tool call is attributed to whichever request opened the session.""" logger, exporter = _logger() - session_opener = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) - this_message = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + session_opener = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + this_message = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(session_opener) trace_token = set_mcp_message_trace_carrier( {"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"} ) transport_token = set_mcp_message_transport_span(this_message) try: - asyncio.run( - logger.async_log_success_event( - {"standard_logging_object": _mcp_list_payload()}, None, None, None - ) - ) + asyncio.run(logger.async_log_success_event({"standard_logging_object": _mcp_list_payload()}, None, None, None)) finally: reset_mcp_message_transport_span(transport_token) reset_mcp_message_trace_carrier(trace_token) @@ -1103,9 +1005,7 @@ def test_pre_call_idempotent_keeps_first_span(): span (with the true start time) is kept, not replaced.""" logger, _ = _logger() kwargs = _kwargs() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) with trace.use_span(server, end_on_exit=False): logger.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs) first = logger._open_llm_calls["call_1"] @@ -1124,9 +1024,7 @@ def test_llm_span_parents_to_ambient_server_span(): """The span is opened at ``pre_call`` while the server span is the active context, so it nests under it natively (no ``litellm_parent_otel_span``).""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) _emit_llm(logger, ambient=server) server.end() by_name = {s.name: s for s in exporter.get_finished_spans()} @@ -1158,9 +1056,7 @@ def test_llm_span_anchors_to_root_even_inside_active_phase_span(): span is the *active* context. The LLM span must still parent to the request root (the server span), never to the auth span it happens to be nested in.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(server) kwargs = _kwargs() # ``auth`` phase span is the active span when pre_call + close run. @@ -1183,9 +1079,7 @@ def test_live_llm_span_anchors_to_root_with_no_active_span(): instead of orphaning — and the detached close just ends it, in the right trace.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(server) kwargs = _kwargs() logger.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs) @@ -1203,9 +1097,7 @@ def test_deferred_llm_span_reads_anchor_at_close(): sync-only provider's thread-pool call) the span defers; the close — back on the request task, anchor visible — must parent it to the root, not orphan it.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) kwargs = _kwargs() # pre_call with NO anchor and no active span → deferred. logger.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs) @@ -1228,9 +1120,7 @@ def test_synthetic_error_log_produces_no_llm_span(): from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(server) payload = _payload( status="failure", @@ -1255,19 +1145,12 @@ def test_create_request_started_span_captures_anchor(): from litellm.integrations.otel.plumbing.context import request_root_span logger, _ = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) with trace.use_span(server, end_on_exit=False): - returned = logger.create_litellm_proxy_request_started_span( - start_time=datetime.now(), headers=None - ) + returned = logger.create_litellm_proxy_request_started_span(start_time=datetime.now(), headers=None) server.end() assert returned.get_span_context().span_id == server.get_span_context().span_id - assert ( - request_root_span().get_span_context().span_id - == server.get_span_context().span_id - ) + assert request_root_span().get_span_context().span_id == server.get_span_context().span_id def test_guardrail_span_anchors_to_root_inside_active_phase_span(): @@ -1275,9 +1158,7 @@ def test_guardrail_span_anchors_to_root_inside_active_phase_span(): span must still be a sibling of the LLM call under the request root, not a child of auth.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(server) entry = {"guardrail_name": "my_guard", "guardrail_status": "success"} with trace.use_span(server, end_on_exit=False): @@ -1315,9 +1196,7 @@ def test_async_post_call_failure_hook_stamps_error_on_root_span(): set_request_root_span(server) exc = _proxy_exc("litellm.BadRequestError: messages is required", 400) result = asyncio.run( - logger.async_post_call_failure_hook( - request_data={}, original_exception=exc, user_api_key_dict=UserAPIKeyAuth() - ) + logger.async_post_call_failure_hook(request_data={}, original_exception=exc, user_api_key_dict=UserAPIKeyAuth()) ) server.end() assert result is None @@ -1478,9 +1357,7 @@ def test_record_error_attributes_on_span_does_not_duplicate_an_already_stamped_e set_request_root_span(server) exc = _proxy_exc("Authentication Error, invalid key", 401) asyncio.run( - logger.async_post_call_failure_hook( - request_data={}, original_exception=exc, user_api_key_dict=UserAPIKeyAuth() - ) + logger.async_post_call_failure_hook(request_data={}, original_exception=exc, user_api_key_dict=UserAPIKeyAuth()) ) logger.record_error_attributes_on_span(server, exc, 400) server.end() @@ -1562,14 +1439,8 @@ def test_real_logging_pre_call_opens_span_end_to_end(): # pre_call fires log_pre_api_call → opens the boundary span on the obj. logging_obj.pre_call(input="hi", api_key="sk-test") # The success callback closes it, reading the typed payload. - logging_obj.model_call_details["standard_logging_object"] = _payload( - litellm_call_id="call_e2e" - ) - asyncio.run( - logger.async_log_success_event( - logging_obj.model_call_details, None, None, None - ) - ) + logging_obj.model_call_details["standard_logging_object"] = _payload(litellm_call_id="call_e2e") + asyncio.run(logger.async_log_success_event(logging_obj.model_call_details, None, None, None)) finally: monkeypatch.undo() (span,) = exporter.get_finished_spans() @@ -1586,9 +1457,7 @@ def test_deferred_span_parents_to_ambient_at_close(): kwargs = _kwargs() # pre_call with NO ambient span (the thread-pool case) → deferred. logger.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs) - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) # The close callback runs with the (worker-copied) server span ambient. with trace.use_span(server, end_on_exit=False): asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) @@ -1647,9 +1516,7 @@ def test_provider_model_and_team_metadata_on_real_boundary_flow(): import json logger, exporter = _logger(team_metadata_keys=["tier", "cost_center"]) - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) payload = _payload( hidden_params={"litellm_model_name": "azure/my-deployment"}, metadata={ @@ -1688,9 +1555,7 @@ def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans(): sibling such as ``requester_ip_address`` is not stamped from here even though the default allowlist names it, and an unlisted caller key is not promoted.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) data = { "model": "gpt-4o", "metadata": {"requester_ip_address": "127.0.0.1", "requester_metadata": {"trace_id": "abc"}}, @@ -1700,9 +1565,7 @@ def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans(): # pre-call seeds baggage + stamps the active server span await logger.async_pre_call_hook(_Auth(), None, data, "completion") # a later service call (same task) must inherit the identity - await logger.async_service_success_hook( - payload=_ServicePayload("redis", "set"), parent_otel_span=server - ) + await logger.async_service_success_hook(payload=_ServicePayload("redis", "set"), parent_otel_span=server) with trace.use_span(server, end_on_exit=False): asyncio.run(_flow()) @@ -1714,9 +1577,7 @@ def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans(): assert redis.attributes[LiteLLM.KEY_HASH] == "hash1" assert redis.attributes[f"{LiteLLM.METADATA_PREFIX}user_api_key_user_id"] == "u1" srv = spans[LITELLM_PROXY_REQUEST_SPAN_NAME] - assert ( - srv.attributes[LiteLLM.TEAM_ID] == "t1" - ) # stamped directly on the server span + assert srv.attributes[LiteLLM.TEAM_ID] == "t1" # stamped directly on the server span assert srv.attributes[f"{LiteLLM.METADATA_PREFIX}user_api_key_user_id"] == "u1" assert not any( k in (f"{LiteLLM.METADATA_PREFIX}requester_ip_address", f"{LiteLLM.METADATA_PREFIX}trace_id") @@ -1773,18 +1634,17 @@ class _Service: class _ServicePayload: - def __init__(self, service="redis", call_type="set", error=None, caller=None): + def __init__(self, service="redis", call_type="set", error=None, caller=None, target=None): self.service = _Service(service) self.call_type = call_type self.caller = caller + self.target = target self.error = error def _service_parent(logger): """Helper: a live PROXY_REQUEST span to parent service spans under.""" - return logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + return logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) async def _redis_get_through_service_logger(logger): @@ -1810,29 +1670,176 @@ async def _redis_get_through_service_logger(logger): MagicMock(get_cache=MagicMock(return_value=None)), ), ): - cache = RedisCache( - host="127.0.0.1", port=6379, service_logger_obj=ServiceLogging() - ) + cache = RedisCache(host="127.0.0.1", port=6379, service_logger_obj=ServiceLogging()) await cache.async_get_cache("otel-naming-key") - await asyncio.gather( - *(t for t in asyncio.all_tasks() if t is not asyncio.current_task()) - ) + await asyncio.gather(*(t for t in asyncio.all_tasks() if t is not asyncio.current_task())) def test_redis_service_span_is_named_by_operation_and_keeps_the_caller_chain_as_an_attribute(): - """``redis async_get_cache``, not ``redis async_get_cache <- caller <- caller``: the stack - walk that used to be spliced into the span name rides on ``litellm.service.caller`` instead, - so one operation is one span name and ``db.operation.name`` is the bare operation.""" + """``redis.get``, not ``redis async_get_cache <- caller <- caller``: the stack walk that + used to be spliced into the span name rides on ``litellm.service.caller`` instead, and the + method name on ``db.operation.name``, so one operation is one span name.""" logger, exporter = _logger() asyncio.run(_redis_get_through_service_logger(logger)) (span,) = [s for s in exporter.get_finished_spans() if s.name.startswith("redis")] - assert span.name == "redis async_get_cache" + assert span.name == "redis.get" assert span.attributes[LiteLLM.SERVICE_CALL_TYPE] == "async_get_cache" assert span.attributes["db.operation.name"] == "async_get_cache" - callers = span.attributes[LiteLLM.SERVICE_CALLER].split(" <- ") - assert callers[0] == "_redis_get_through_service_logger" and len(callers) == 2, ( - callers + assert span.attributes[LiteLLM.SERVICE_CALLER] == "_redis_get_through_service_logger" + assert LiteLLM.SERVICE_TARGET not in span.attributes + + +def test_service_span_is_named_by_purpose_when_the_producer_declares_a_target(): + """``redis.get llm_response``, the ``{operation} {target}`` shape the OTel database + conventions ask for, while the raw method name stays on the attributes dashboards + filter on (``litellm.service.call_type``, ``db.operation.name`` and the V1 ``call_type``).""" + logger, exporter = _logger() + parent = _service_parent(logger) + try: + asyncio.run( + logger.async_service_success_hook( + payload=_ServicePayload( + "redis", + "async_get_cache", + caller="_retrieve_from_cache <- _async_get_cache", + target="llm_response", + ), + parent_otel_span=parent, + ) + ) + finally: + parent.end() + (span,) = [s for s in exporter.get_finished_spans() if s.name.startswith("redis")] + assert span.name == "redis.get llm_response" + assert span.kind is SpanKind.CLIENT + assert span.attributes[LiteLLM.SERVICE_CALL_TYPE] == "async_get_cache" + assert span.attributes["db.operation.name"] == "async_get_cache" + assert span.attributes["call_type"] == "async_get_cache" + assert span.attributes[LiteLLM.SERVICE_TARGET] == "llm_response" + assert span.attributes[LiteLLM.SERVICE_CALLER] == "_retrieve_from_cache <- _async_get_cache" + + +@pytest.mark.parametrize( + ("call_type", "targeted", "untargeted"), + [ + ("async_get_cache", "redis.get auth_objects", "redis.get"), + ("async_batch_get_cache", "redis.mget auth_objects", "redis.mget"), + ("async_set_cache_pipeline_with_ttls", "redis.set auth_objects", "redis.set"), + ("async_increment_pipeline", "redis.incr auth_objects", "redis.incr"), + ("async_delete_cache", "redis.delete auth_objects", "redis.delete"), + ("async_scan_iter", "redis.scan auth_objects", "redis.scan"), + ("request_redis_batch", "redis.pipeline auth_objects", "redis.pipeline"), + ("async_frobnicate", "redis async_frobnicate", "redis async_frobnicate"), + ], +) +def test_service_span_verb_follows_the_cache_method_behind_the_call(call_type, targeted, untargeted): + """Every known Redis method renders as ``redis.{verb}``, with the key family appended when + the producer declared one, so one trace never mixes ``redis.get llm_response`` with + ``redis async_get_cache``; an unknown method keeps the raw ``{service} {call_type}`` name.""" + from litellm.integrations.otel.model.payloads import ServiceSpanData + from litellm.integrations.otel.model.spans import service_span_name + + assert ( + service_span_name(ServiceSpanData(service_name="redis", call_type=call_type, target="auth_objects")) == targeted ) + assert service_span_name(ServiceSpanData(service_name="redis", call_type=call_type)) == untargeted + + +def test_postgres_service_span_keeps_its_function_name_inside_a_targeted_phase(): + """A DB helper that runs inside ``service_target("auth_objects")`` (the whole auth phase + does) is still ``postgres get_data``: the verb scheme is for cache methods, the Postgres + rename to ``db.select {table}`` is a separate change.""" + from litellm.integrations.otel.model.payloads import ServiceSpanData + from litellm.integrations.otel.model.spans import service_span_name + + data = ServiceSpanData(service_name="postgres", call_type="get_data", target="auth_objects") + assert service_span_name(data) == "postgres get_data" + + +def test_service_target_declared_by_the_producer_rides_the_service_logger_payload(): + """``service_target`` is a contextvar the real ``ServiceLogging`` stamps onto the payload, + so a producer names its key family once and every cache read inside picks it up.""" + logger, exporter = _logger() + + async def _lookup(): + with service_target("llm_response"): + await _redis_get_through_service_logger(logger) + + asyncio.run(_lookup()) + (span,) = [s for s in exporter.get_finished_spans() if s.name.startswith("redis")] + assert span.name == "redis.get llm_response" + assert span.attributes[LiteLLM.SERVICE_TARGET] == "llm_response" + assert span.attributes[LiteLLM.SERVICE_CALL_TYPE] == "async_get_cache" + + +def test_response_cache_lookup_nests_its_redis_read_under_a_cache_get_span_on_the_request_root(): + """The lookup runs inside a live ``cache.get llm_response`` phase span, a child of the + server span, so the Redis GET is its child and sits before ``chat {model}`` in causal order + instead of landing flat on the root.""" + logger, exporter = _logger() + server = _service_parent(logger) + + async def _lookup(): + with logger.start_phase_span("cache.get llm_response"): + await logger.async_service_success_hook( + payload=_ServicePayload("redis", "async_get_cache", target="llm_response"), + parent_otel_span=server, + ) + + try: + with trace.use_span(server, end_on_exit=False): + asyncio.run(_lookup()) + finally: + server.end() + by_name = {s.name: s for s in exporter.get_finished_spans()} + phase = by_name["cache.get llm_response"] + redis = by_name["redis.get llm_response"] + request_ctx = server.get_span_context() + assert phase.kind is SpanKind.INTERNAL + assert phase.parent.span_id == request_ctx.span_id + assert not phase.links + assert redis.parent.span_id == phase.context.span_id + assert redis.context.trace_id == request_ctx.trace_id + assert not redis.links + + +def test_response_cache_write_from_the_post_response_phase_is_one_linked_trace(): + """The write runs after the response is on the wire, so its ``cache.set llm_response`` + span detaches from the request as a linked root (the request trace keeps its real + duration), and the Redis SET it issues nests under that root instead of detaching + into a third, unrelated trace.""" + logger, exporter = _logger() + server = _service_parent(logger) + + async def _write_task(): + with logger.start_phase_span("cache.set llm_response"): + await logger.async_service_success_hook( + payload=_ServicePayload("redis", "async_set_cache", target="llm_response"), + parent_otel_span=server, + ) + + async def _request(): + with post_response_phase(): + task = asyncio.create_task(_write_task()) + await task + + try: + with trace.use_span(server, end_on_exit=False): + asyncio.run(_request()) + finally: + server.end() + by_name = {s.name: s for s in exporter.get_finished_spans()} + phase = by_name["cache.set llm_response"] + redis = by_name["redis.set llm_response"] + request_ctx = server.get_span_context() + assert phase.parent is None + assert phase.context.trace_id != request_ctx.trace_id + assert [(link.context.trace_id, link.context.span_id) for link in phase.links] == [ + (request_ctx.trace_id, request_ctx.span_id) + ] + assert redis.parent.span_id == phase.context.span_id + assert redis.context.trace_id == phase.context.trace_id + assert not redis.links def test_async_service_success_hook_emits_service_span(): @@ -1946,11 +1953,7 @@ def test_metrics_only_ping_without_timing_or_parent_is_noop(): per-request ``self`` latency hook, in-memory queue gauges) — not a traceable operation, so no span is emitted.""" logger, exporter = _logger() - asyncio.run( - logger.async_service_success_hook( - payload=_ServicePayload(), parent_otel_span=None - ) - ) + asyncio.run(logger.async_service_success_hook(payload=_ServicePayload(), parent_otel_span=None)) assert exporter.get_finished_spans() == () @@ -2009,22 +2012,13 @@ def test_metrics_only_services_emit_no_span(): def test_service_span_inherits_parent_when_provided(): logger, exporter = _logger() - parent = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + parent = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) try: - asyncio.run( - logger.async_service_success_hook( - payload=_ServicePayload(), parent_otel_span=parent - ) - ) + asyncio.run(logger.async_service_success_hook(payload=_ServicePayload(), parent_otel_span=parent)) finally: parent.end() by_name = {s.name: s for s in exporter.get_finished_spans()} - assert ( - by_name["redis set"].parent.span_id - == by_name[LITELLM_PROXY_REQUEST_SPAN_NAME].get_span_context().span_id - ) + assert by_name["redis set"].parent.span_id == by_name[LITELLM_PROXY_REQUEST_SPAN_NAME].get_span_context().span_id def test_service_span_prefers_ambient_context_over_threaded_parent(): @@ -2034,9 +2028,7 @@ def test_service_span_prefers_ambient_context_over_threaded_parent(): ambient has no live span (a background service call).""" logger, exporter = _logger() ambient = logger._emitter.start_span(SpanRole.LLM_CALL, "chat gpt-4o") - threaded = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + threaded = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) try: with trace.use_span(ambient, end_on_exit=False): asyncio.run( @@ -2137,9 +2129,7 @@ def test_service_call_under_a_remote_parent_is_never_detached(): assert list(span.links) == [] -def _service_hook_from_post_response_task( - logger, payload, *, parent, ambient, end_time -): +def _service_hook_from_post_response_task(logger, payload, *, parent, ambient, end_time): """Log ``payload`` the way the proxy's post-response tail does: the hook runs on a task spawned from inside ``post_response_phase`` while the server span is still open.""" @@ -2153,9 +2143,7 @@ def _service_hook_from_post_response_task( end_time=end_time, ) ) - assert not in_post_response_phase(), ( - "the phase must not leak into the request task" - ) + assert not in_post_response_phase(), "the phase must not leak into the request task" await task if ambient is None: @@ -2174,9 +2162,7 @@ def test_service_call_from_the_post_response_phase_detaches_before_the_server_sp so the call ends before its parent does. Timing alone would keep it a child; being dispatched from the post-response phase is what detaches it, with a link.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) assert server.is_recording() try: _service_hook_from_post_response_task( @@ -2188,7 +2174,7 @@ def test_service_call_from_the_post_response_phase_detaches_before_the_server_sp ) finally: server.end(end_time=to_ns(_REQUEST_END)) - span = {s.name: s for s in exporter.get_finished_spans()}["redis async_set_cache"] + span = {s.name: s for s in exporter.get_finished_spans()}["redis.set"] request_ctx = server.get_span_context() assert span.end_time < server.end_time assert span.parent is None @@ -2237,7 +2223,9 @@ def test_redis_write_from_a_success_callback_detaches_while_the_server_span_is_s class _RedisWritingCallback(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): await logger.async_service_success_hook( - payload=_ServicePayload("redis", "async_increment", caller="async_increment_cache <- async_log_success_event"), + payload=_ServicePayload( + "redis", "async_increment", caller="async_increment_cache <- async_log_success_event" + ), parent_otel_span=None, start_time=_REQUEST_END - 0.5, end_time=_REQUEST_END - 0.1, @@ -2273,7 +2261,7 @@ def test_redis_write_from_a_success_callback_detaches_while_the_server_span_is_s asyncio.run(_request()) finally: server.end(end_time=to_ns(_REQUEST_END)) - span = {s.name: s for s in exporter.get_finished_spans()}["redis async_increment"] + span = {s.name: s for s in exporter.get_finished_spans()}["redis.incr"] request_ctx = server.get_span_context() assert span.end_time < server.end_time assert span.parent is None @@ -2303,13 +2291,9 @@ def test_create_proxy_request_started_span_returns_ambient_span(): ) assert exporter.get_finished_spans() == () # With an active server span, return it (do NOT create a new one). - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) with trace.use_span(server, end_on_exit=False): - got = logger.create_litellm_proxy_request_started_span( - start_time=datetime.now(timezone.utc), headers=None - ) + got = logger.create_litellm_proxy_request_started_span(start_time=datetime.now(timezone.utc), headers=None) server.end() assert got is server @@ -2341,9 +2325,7 @@ def test_default_config_reads_env(monkeypatch): monkeypatch.delenv("OTEL_EXPORTER", raising=False) monkeypatch.delenv("OTEL_EXPORTER_OTLP_PROTOCOL", raising=False) logger = OpenTelemetryV2( - tracer_provider=providers.build_tracer_provider( - OpenTelemetryV2Config(exporter="in_memory") - ) + tracer_provider=providers.build_tracer_provider(OpenTelemetryV2Config(exporter="in_memory")) ) assert logger.config.exporter == "console" @@ -2381,9 +2363,7 @@ def test_select_global_otel_v2_logger_reuses_existing_preset_logger(): cfg = OpenTelemetryV2Config(exporter="in_memory") tp = providers.build_tracer_provider(cfg) - preset_logger = OpenTelemetryV2( - config=cfg, callback_name="arize", tracer_provider=tp - ) + preset_logger = OpenTelemetryV2(config=cfg, callback_name="arize", tracer_provider=tp) chosen = select_global_otel_v2_logger([object(), preset_logger, object()]) assert chosen is preset_logger @@ -2444,14 +2424,10 @@ def test_publish_global_otel_v2_provider_sets_selected_logger_provider(monkeypat monkeypatch.setattr(otel_logger, "_published_v2_provider", None) cfg = OpenTelemetryV2Config(exporter="in_memory") tp = providers.build_tracer_provider(cfg) - preset_logger = OpenTelemetryV2( - config=cfg, callback_name="arize", tracer_provider=tp - ) + preset_logger = OpenTelemetryV2(config=cfg, callback_name="arize", tracer_provider=tp) published = [] - chosen = publish_global_otel_v2_provider( - [object(), preset_logger], published.append - ) + chosen = publish_global_otel_v2_provider([object(), preset_logger], published.append) assert chosen is preset_logger assert published == [preset_logger._tracer_provider] @@ -2475,9 +2451,7 @@ def test_registers_into_litellm_service_callback(monkeypatch): # A second OTel logger sees one is already registered and does not duplicate. OpenTelemetryV2(config=cfg, tracer_provider=tp) otel_registrations = [ - cb - for cb in litellm.service_callback - if cb.__class__.__module__.startswith("litellm.integrations.otel") + cb for cb in litellm.service_callback if cb.__class__.__module__.startswith("litellm.integrations.otel") ] assert len(otel_registrations) == 1 @@ -2500,9 +2474,7 @@ def test_registers_into_litellm_input_callback(monkeypatch): OpenTelemetryV2(config=cfg, tracer_provider=tp) otel_registrations = [ - cb - for cb in litellm.input_callback - if cb.__class__.__module__.startswith("litellm.integrations.otel") + cb for cb in litellm.input_callback if cb.__class__.__module__.startswith("litellm.integrations.otel") ] assert len(otel_registrations) == 1 @@ -2540,9 +2512,7 @@ def test_registers_into_async_success_and_failure_callbacks(monkeypatch): litellm._async_failure_callback, ): otel_registrations = [ - cb - for cb in callback_list - if cb.__class__.__module__.startswith("litellm.integrations.otel") + cb for cb in callback_list if cb.__class__.__module__.startswith("litellm.integrations.otel") ] assert len(otel_registrations) == 1 @@ -2589,14 +2559,8 @@ def test_boundary_span_closes_without_proxy_fanout(monkeypatch): assert "pt_leak" in logger._open_llm_calls # The close runs through the real async_success_handler, which iterates # _async_success_callback — where the logger self-registered. - logging_obj.model_call_details["standard_logging_object"] = _payload( - litellm_call_id="pt_leak" - ) - asyncio.run( - logging_obj.async_success_handler( - result=None, start_time=datetime.now(), end_time=datetime.now() - ) - ) + logging_obj.model_call_details["standard_logging_object"] = _payload(litellm_call_id="pt_leak") + asyncio.run(logging_obj.async_success_handler(result=None, start_time=datetime.now(), end_time=datetime.now())) assert "pt_leak" not in logger._open_llm_calls # carrier closed, not leaked (span,) = exporter.get_finished_spans() assert span.name == "chat gpt-4o" @@ -2623,18 +2587,14 @@ def test_guardrail_span_parents_to_ambient_server_span(): ambient, so with no explicit anchor set the guardrail span parents to it. (Auth already finished, so no phase span is active.)""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) entry = _guardrail_entry(start=1000.0, end=1000.5) try: with trace.use_span(server, end_on_exit=False): logger.emit_guardrail_span(entry) finally: server.end() - g = {s.name: s for s in exporter.get_finished_spans()}[ - "execute_guardrail openai-moderation" - ] + g = {s.name: s for s in exporter.get_finished_spans()}["execute_guardrail openai-moderation"] assert g.parent.span_id == server.get_span_context().span_id @@ -2642,18 +2602,14 @@ def test_guardrail_span_uses_actual_execution_timestamps(): """A pre_call guardrail's span carries its real start/end (from the logging entry), so it sorts before the LLM call instead of at emission time.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) entry = _guardrail_entry(start=1700.0, end=1700.25) try: with trace.use_span(server, end_on_exit=False): logger.emit_guardrail_span(entry) finally: server.end() - g = {s.name: s for s in exporter.get_finished_spans()}[ - "execute_guardrail openai-moderation" - ] + g = {s.name: s for s in exporter.get_finished_spans()}["execute_guardrail openai-moderation"] assert g.start_time == to_ns(1700.0) assert g.end_time == to_ns(1700.25) @@ -2664,9 +2620,7 @@ def test_emit_guardrail_span_anchors_to_root_not_ambient_phase_span(): ambient, so a guardrail emitted mid-``auth`` is a sibling of the LLM call, not a child of ``auth``.""" logger, exporter = _logger() - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) set_request_root_span(server) entry = _guardrail_entry(start=2000.0, end=2000.1) with logger.start_phase_span("auth /chat/completions"): @@ -2725,11 +2679,7 @@ def _emitted_metric_names(reader) -> set: if data is None: return set() return { - m.name - for rm in data.resource_metrics - for sm in rm.scope_metrics - for m in sm.metrics - if any(m.data.data_points) + m.name for rm in data.resource_metrics for sm in rm.scope_metrics for m in sm.metrics if any(m.data.data_points) } @@ -2786,23 +2736,11 @@ def test_invalid_metric_filter_logged_once_records_nothing(caplog, monkeypatch): with caplog.at_level(logging.ERROR, logger="LiteLLM"): # Neither call may raise; the bad filter is caught in the logger. - asyncio.run( - logger.async_log_success_event( - _metric_success_kwargs(), response_obj, start, end - ) - ) - asyncio.run( - logger.async_log_success_event( - _metric_success_kwargs(), response_obj, start, end - ) - ) + asyncio.run(logger.async_log_success_event(_metric_success_kwargs(), response_obj, start, end)) + asyncio.run(logger.async_log_success_event(_metric_success_kwargs(), response_obj, start, end)) assert _emitted_metric_names(reader) == set() # nothing recorded - errors = [ - r - for r in caplog.records - if r.levelno == logging.ERROR and "metric filter" in r.getMessage() - ] + errors = [r for r in caplog.records if r.levelno == logging.ERROR and "metric filter" in r.getMessage()] assert len(errors) == 1 # logged once, second bad record does not re-log @@ -2919,9 +2857,7 @@ def _phoenix_routing_logger(capture_kind): ) default_exporter = InMemorySpanExporter() tracer_provider = providers.build_tracer_provider(cfg, exporter=default_exporter) - logger = OpenTelemetryV2( - config=cfg, callback_name="arize_phoenix", tracer_provider=tracer_provider - ) + logger = OpenTelemetryV2(config=cfg, callback_name="arize_phoenix", tracer_provider=tracer_provider) return logger, default_exporter, captured @@ -2971,9 +2907,7 @@ def test_project_routing_resolves_at_pre_call_before_payload_exists(): logger, default_exporter, captured = _phoenix_routing_logger("capture_route_c") auth_md = {"phoenix_project_name": "team-proj"} litellm_params = {"metadata": {"user_api_key_auth_metadata": auth_md}} - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) with trace.use_span(server, end_on_exit=False): logger.log_pre_api_call( model="gpt-4o", @@ -2984,9 +2918,7 @@ def test_project_routing_resolves_at_pre_call_before_payload_exists(): assert len(captured) == 1 # routed exporter already built at pre_call close_kwargs = { - "standard_logging_object": _payload( - metadata={"user_api_key_auth_metadata": auth_md} - ), + "standard_logging_object": _payload(metadata={"user_api_key_auth_metadata": auth_md}), "litellm_params": litellm_params, } asyncio.run(logger.async_log_success_event(close_kwargs, None, None, None)) @@ -2999,9 +2931,7 @@ def test_project_routing_resolves_at_pre_call_before_payload_exists(): assert routed_span.parent is None (link,) = routed_span.links assert link.context.span_id == server.get_span_context().span_id - assert all( - s.name != "chat gpt-4o" for s in default_exporter.get_finished_spans() - ) + assert all(s.name != "chat gpt-4o" for s in default_exporter.get_finished_spans()) def test_evicted_provider_still_exports_span_opened_before_eviction(monkeypatch): @@ -3013,9 +2943,7 @@ def test_evicted_provider_still_exports_span_opened_before_eviction(monkeypatch) monkeypatch.setattr(routing_mod, "_MAX_CACHED_PROVIDERS", 1) logger, _default_exporter, captured = _phoenix_routing_logger("capture_evict") md_a = {"user_api_key_auth_metadata": {"phoenix_project_name": "proj-a"}} - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) with trace.use_span(server, end_on_exit=False): logger.log_pre_api_call( model="gpt-4o", @@ -3066,9 +2994,7 @@ def test_deferred_pre_call_does_not_churn_tenant_cache(monkeypatch): monkeypatch.setattr(routing_mod, "_shutdown_provider", lambda p: shut_down.append(p)) logger, _default, captured = _phoenix_routing_logger("capture_deferred_churn") md_a = {"user_api_key_auth_metadata": {"phoenix_project_name": "proj-a"}} - server = logger._emitter.start_span( - SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME - ) + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) with trace.use_span(server, end_on_exit=False): _emit_llm( logger, @@ -3132,9 +3058,7 @@ def test_success_without_pre_call_emits_deferred_span(): logger, exporter = _logger() # No log_pre_api_call: this logger never receives the input hook. The # request-level provider-handoff stamp is present (pre_call ran globally). - asyncio.run( - logger.async_log_success_event({**_kwargs(), "api_call_start_time": 100.0}, None, 100.0, 101.5) - ) + asyncio.run(logger.async_log_success_event({**_kwargs(), "api_call_start_time": 100.0}, None, 100.0, 101.5)) spans = exporter.get_finished_spans() assert len(spans) == 1 assert spans[0].attributes.get("gen_ai.operation.name") @@ -3174,9 +3098,7 @@ def test_failure_without_pre_call_emits_deferred_error_span(): error_information={"error_class": "RateLimitError", "error_code": "429"}, ) asyncio.run( - logger.async_log_failure_event( - {**_kwargs(payload=payload), "api_call_start_time": 100.0}, None, None, None - ) + logger.async_log_failure_event({**_kwargs(payload=payload), "api_call_start_time": 100.0}, None, None, None) ) spans = exporter.get_finished_spans() assert len(spans) == 1 @@ -3238,3 +3160,165 @@ def test_provisional_close_then_payload_close_does_not_duplicate(): server.end() llm_spans = [s for s in exporter.get_finished_spans() if s.name.startswith("chat")] assert len(llm_spans) == 1 + + +def test_pipeline_op_count_lands_as_an_int_on_both_metadata_keys(): + """A ``RedisBatch`` flush reports ``call_type=request_redis_batch`` with + ``event_metadata={"op_count": N}``; the span is ``redis.pipeline`` (no ``[N]`` in the + name) and the count survives sanitization as an int on the namespaced V2 key and the + bare V1 key, so a dashboard can sum it.""" + logger, exporter = _logger() + parent = _service_parent(logger) + try: + asyncio.run( + logger.async_service_success_hook( + payload=_ServicePayload("redis", "request_redis_batch"), + parent_otel_span=parent, + event_metadata={"op_count": 3}, + ) + ) + finally: + parent.end() + (span,) = [s for s in exporter.get_finished_spans() if s.name.startswith("redis")] + assert span.name == "redis.pipeline" + assert span.attributes[LiteLLM.SERVICE_CALL_TYPE] == "request_redis_batch" + v2_count = span.attributes[f"{LiteLLM.METADATA_PREFIX}op_count"] + v1_count = span.attributes["op_count"] + assert (v2_count, v1_count) == (3, 3) + assert type(v2_count) is int and type(v1_count) is int + + +_REDIS_CACHE_MODULES = ( + "litellm/caching/redis_cache.py", + "litellm/caching/redis_cluster_cache.py", + "litellm/caching/redis_semantic_cache.py", + "litellm/caching/dual_cache.py", + "litellm/caching/redis_batch.py", +) + + +def _redis_call_types_emitted_by_the_cache_layer(): + """Every ``call_type`` literal the Redis cache layer hands to the service logger, plus the + two ``RedisBatch`` names it passes as ``call_type=self.name``.""" + import re + from pathlib import Path + + repo = Path(__file__).resolve().parents[4] + sources = "\n".join((repo / module).read_text() for module in _REDIS_CACHE_MODULES) + literal = frozenset(re.findall(r'call_type="([a-z_]+)"', sources)) + batch_names = frozenset(re.findall(r'RedisBatch\([^)]*name="([a-z_]+)"', sources)) + return sorted(literal | batch_names) + + +def test_no_call_type_the_redis_cache_layer_emits_can_fall_back_to_the_raw_method_name(): + """The ``{service} {call_type}`` branch exists for services without a verb scheme; for Redis + it must be unreachable, or one trace mixes ``redis.get llm_response`` with ``redis async_get_cache`` + again the moment a cache method is added without a verb.""" + from litellm.integrations.otel.model.payloads import ServiceSpanData + from litellm.integrations.otel.model.spans import service_span_name + + call_types = _redis_call_types_emitted_by_the_cache_layer() + assert {"async_get_cache", "async_batch_get_cache", "request_redis_batch", "post_call_redis_batch"} <= set( + call_types + ) + fallbacks = [ + call_type + for call_type in call_types + if not service_span_name(ServiceSpanData(service_name="redis", call_type=call_type)).startswith("redis.") + ] + assert fallbacks == [] + + +def test_mixed_pipeline_families_land_on_their_own_bounded_attribute(): + """A flush that carried several owners' ops is ``redis.pipeline mixed``; the sorted family + list goes to ``litellm.redis.families`` (not under ``litellm.metadata.``) while ``op_count`` + stays an int on ``litellm.metadata.op_count``, and the legacy vocabulary keeps the bare keys.""" + logger, exporter = _logger() + parent = _service_parent(logger) + try: + asyncio.run( + logger.async_service_success_hook( + payload=_ServicePayload("redis", "request_redis_batch", target="mixed"), + parent_otel_span=parent, + event_metadata={"op_count": 3, "families": "auth_objects,spend_counters"}, + ) + ) + finally: + parent.end() + (span,) = [s for s in exporter.get_finished_spans() if s.name.startswith("redis")] + assert span.name == "redis.pipeline mixed" + assert span.attributes[LiteLLM.REDIS_FAMILIES] == "auth_objects,spend_counters" + assert span.attributes["families"] == "auth_objects,spend_counters" + assert f"{LiteLLM.METADATA_PREFIX}families" not in span.attributes + assert span.attributes[f"{LiteLLM.METADATA_PREFIX}op_count"] == 3 + assert type(span.attributes[f"{LiteLLM.METADATA_PREFIX}op_count"]) is int + + +def test_deferred_close_starts_at_the_provider_handoff_not_the_logging_objects_birth(): + """A destination logger never sees ``pre_call``, so its copy of ``chat`` is created at + close. Starting it at the logging object's ``start_time`` makes it span routing and the + cache lookup; the provider handoff stamp is where the attempt really began.""" + logger, exporter = _logger() + handoff = datetime(2026, 5, 26, 12, 0, 0, 500000, tzinfo=timezone.utc) + logging_start = datetime(2026, 5, 26, 12, 0, 0, tzinfo=timezone.utc) + kwargs = {**_kwargs(), "api_call_start_time": handoff} + + asyncio.run(logger.async_log_success_event(kwargs, None, logging_start, None)) + + (span,) = exporter.get_finished_spans() + assert span.name == "chat gpt-4o" + assert span.start_time == to_ns(handoff) + + +def test_a_call_made_inside_a_phase_nests_under_it_and_names_its_purpose(): + """The auto-router classifier is a chat call litellm makes while picking a deployment. It + belongs under ``route {model_group}``, not beside the caller's own ``chat``, and carries a + bounded purpose so the two are told apart without reading model names.""" + logger, exporter = _logger() + root = logger.tracer.start_span("POST /v1/chat/completions", kind=SpanKind.SERVER) + set_request_root_span(root) + classifier_kwargs = _kwargs(_payload(model="gpt-4o-mini")) + classifier_kwargs["litellm_params"]["metadata"][INTERNAL_CALL_ORIGIN_METADATA_KEY] = ( + AUTOROUTER_CLASSIFIER_CALL_ORIGIN + ) + + with trace.use_span(root, end_on_exit=True): + with logger.start_phase_span("route auto-router") as route: + _emit_llm(logger, classifier_kwargs) + _emit_llm(logger, _kwargs()) + + by_name = {span.name: span for span in exporter.get_finished_spans()} + classifier = by_name["chat gpt-4o-mini"] + assert classifier.parent.span_id == route.get_span_context().span_id + assert classifier.attributes[LiteLLM.REQUEST_PURPOSE] == AUTOROUTER_CLASSIFIER_CALL_ORIGIN + assert by_name["route auto-router"].parent.span_id == root.get_span_context().span_id + provider_call = by_name["chat gpt-4o"] + assert provider_call.parent.span_id == root.get_span_context().span_id + assert LiteLLM.REQUEST_PURPOSE not in provider_call.attributes + + +def test_a_classifier_closed_without_a_carrier_still_nests_under_the_route_phase(): + """A key or team destination logger is a success callback only, so it creates the classifier's + span at close. That span belongs under ``route {model_group}`` in the destination's trace just as + it does in the operator's, not beside the routing it was part of.""" + logger, exporter = _logger() + root = logger.tracer.start_span("POST /v1/chat/completions", kind=SpanKind.SERVER) + set_request_root_span(root) + handoff = datetime(2026, 5, 26, 12, 0, 0, tzinfo=timezone.utc) + classifier_kwargs = { + **_kwargs(_payload(model="gpt-4o-mini", litellm_call_id="call_classifier")), + "api_call_start_time": handoff, + } + classifier_kwargs["litellm_params"]["metadata"][INTERNAL_CALL_ORIGIN_METADATA_KEY] = ( + AUTOROUTER_CLASSIFIER_CALL_ORIGIN + ) + + with trace.use_span(root, end_on_exit=True): + with logger.start_phase_span("route auto-router") as route: + asyncio.run(logger.async_log_success_event(classifier_kwargs, None, None, None)) + asyncio.run(logger.async_log_success_event({**_kwargs(), "api_call_start_time": handoff}, None, None, None)) + + by_name = {span.name: span for span in exporter.get_finished_spans()} + assert by_name["chat gpt-4o-mini"].parent.span_id == route.get_span_context().span_id + assert by_name["chat gpt-4o-mini"].attributes[LiteLLM.REQUEST_PURPOSE] == AUTOROUTER_CLASSIFIER_CALL_ORIGIN + assert by_name["chat gpt-4o"].parent.span_id == root.get_span_context().span_id diff --git a/tests/unit/proxy/hooks/test_sensitive_data_routing.py b/tests/unit/proxy/hooks/test_sensitive_data_routing.py index 463d3c7ef5e..48a42e42c62 100644 --- a/tests/unit/proxy/hooks/test_sensitive_data_routing.py +++ b/tests/unit/proxy/hooks/test_sensitive_data_routing.py @@ -6,9 +6,9 @@ This feature allows guardrails to route requests to a different model All subsequent requests in the same session are routed to the same model. """ -import logging import asyncio -from typing import Any, Dict, Optional +import logging +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -21,22 +21,22 @@ from litellm.integrations.custom_guardrail import ( get_session_id_from_request_data, ) from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.utils import InternalUsageCache from litellm.proxy.hooks.sensitive_data_routing import ( - _PROXY_SensitiveDataRoutingHandler, - SENSITIVE_ROUTING_CACHE_PREFIX, DEFAULT_SENSITIVE_ROUTING_TTL, + SENSITIVE_ROUTING_CACHE_PREFIX, + _PROXY_SensitiveDataRoutingHandler, ) +from litellm.proxy.utils import InternalUsageCache class MockInternalUsageCache: def __init__(self): - self._cache: Dict[str, Any] = {} - self._ttls: Dict[str, int] = {} + self._cache: dict[str, Any] = {} + self._ttls: dict[str, int] = {} self.dual_cache = MagicMock() self.dual_cache.redis_cache = None - async def async_get_cache(self, key: str, **kwargs) -> Optional[Any]: + async def async_get_cache(self, key: str, **kwargs) -> Any | None: return self._cache.get(key) async def async_set_cache(self, key: str, value: Any, ttl: int = 3600, **kwargs): @@ -67,6 +67,37 @@ class TestSensitiveDataRoutingHandler: routed_model = await handler._get_routed_model("test-session-123", key) assert routed_model == "on-premise-model" + @pytest.mark.asyncio + async def test_session_pin_reads_and_writes_are_targeted_as_sensitive_route_pins(self, user_api_key_dict): + """The pin read on every request used to surface as a bare ``redis.get`` in the trace; the hook + declares its key family so the span reads ``redis.get sensitive_route_pins``.""" + from litellm._internal_context import current_service_target + + class TargetRecordingCache(MockInternalUsageCache): + def __init__(self): + super().__init__() + self.targets: list[str | None] = [] + + async def async_get_cache(self, key: str, **kwargs): + self.targets.append(current_service_target()) + return await super().async_get_cache(key, **kwargs) + + async def async_set_cache(self, key: str, value: Any, ttl: int = 3600, **kwargs): + self.targets.append(current_service_target()) + await super().async_set_cache(key, value, ttl=ttl, **kwargs) + + cache = TargetRecordingCache() + handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + await handler.set_session_routing( + session_id="s-1", model="on-premise-model", user_api_key_dict=user_api_key_dict, guardrail_name="g" + ) + data = {"model": "cloud-model", "litellm_session_id": "s-1"} + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, cache=MagicMock(), data=data, call_type="completion" + ) + assert cache.targets == ["sensitive_route_pins"] * len(cache.targets) and len(cache.targets) >= 2 + assert current_service_target() is None + def test_get_session_id_from_metadata(self): data = {"metadata": {"session_id": "session-from-metadata"}} session_id = get_session_id_from_request_data(data) @@ -95,9 +126,7 @@ class TestSensitiveDataRoutingHandler: assert data["model"] == "gpt-4" @pytest.mark.asyncio - async def test_pre_call_hook_with_routing_override( - self, handler, user_api_key_dict - ): + async def test_pre_call_hook_with_routing_override(self, handler, user_api_key_dict): await handler.set_session_routing( session_id="routed-session", model="on-premise-model", @@ -230,7 +259,7 @@ class TestCustomGuardrailSensitiveDataRouting: request_data = {"model": "gpt-4"} - with pytest.raises(ValueError, match='Cannot route sensitive data without a session_id\\. Ensure') as exc_info: + with pytest.raises(ValueError, match="Cannot route sensitive data without a session_id\\. Ensure") as exc_info: guardrail.raise_sensitive_data_route_exception( route_to_model="on-premise-model", request_data=request_data, @@ -320,10 +349,7 @@ class TestStickySessionRouting: assert result is not None assert result["model"] == "on-premise-model" - assert ( - result["metadata"]["sensitive_data_routing_original_model"] - == f"gpt-{i}" - ) + assert result["metadata"]["sensitive_data_routing_original_model"] == f"gpt-{i}" @pytest.mark.asyncio async def test_different_sessions_independent(self, handler, user_api_key_dict): @@ -434,22 +460,13 @@ class TestCacheKeyAndTTL: assert tenant == "user:alice|team:t1|org:o1" def test_resolve_tenant_distinguishes_keyless_principals(self): - tenant_a = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( - UserAPIKeyAuth(api_key=None, user_id="alice") - ) - tenant_b = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( - UserAPIKeyAuth(api_key=None, user_id="bob") - ) + tenant_a = _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="alice")) + tenant_b = _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None, user_id="bob")) assert tenant_a != tenant_b def test_resolve_tenant_defaults_when_anonymous(self): assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(None) == "default" - assert ( - _PROXY_SensitiveDataRoutingHandler._resolve_tenant( - UserAPIKeyAuth(api_key=None) - ) - == "default" - ) + assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(UserAPIKeyAuth(api_key=None)) == "default" class TestCustomGuardrailSessionIdExtraction: @@ -521,10 +538,7 @@ class TestSensitiveDataRouteExceptionStr: session_id="test-session", guardrail_name="pii-detector", ) - assert ( - str(exc) - == "Sensitive data detected by pii-detector. Routing to model: on-premise-model" - ) + assert str(exc) == "Sensitive data detected by pii-detector. Routing to model: on-premise-model" def test_exception_custom_message(self): exc = SensitiveDataRouteException( @@ -550,29 +564,21 @@ class TestRedisCache: handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( return_value="redis-model" ) - result = await handler_with_redis._get_routed_model( - "session-123", UserAPIKeyAuth(api_key="hashed-key") - ) + result = await handler_with_redis._get_routed_model("session-123", UserAPIKeyAuth(api_key="hashed-key")) assert result == "redis-model" @pytest.mark.asyncio - async def test_get_routed_model_backfills_in_memory_after_redis_hit( - self, handler_with_redis - ): + async def test_get_routed_model_backfills_in_memory_after_redis_hit(self, handler_with_redis): cache_key = "{sensitive_route:hashed-key:session-123}:model" key = UserAPIKeyAuth(api_key="hashed-key") handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( return_value="on-premise-model" ) - handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = ( - AsyncMock(return_value=120) - ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = AsyncMock(return_value=120) first = await handler_with_redis._get_routed_model("session-123", key) assert first == "on-premise-model" - assert handler_with_redis.internal_usage_cache._cache[cache_key] == ( - "on-premise-model" - ) + assert handler_with_redis.internal_usage_cache._cache[cache_key] == ("on-premise-model") handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( side_effect=Exception("Redis went down") @@ -587,52 +593,39 @@ class TestRedisCache: handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( return_value="on-premise-model" ) - handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = ( - AsyncMock(return_value=42) - ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = AsyncMock(return_value=42) await handler_with_redis._get_routed_model("session-123", key) assert handler_with_redis.internal_usage_cache._ttls[cache_key] == 42 @pytest.mark.asyncio - async def test_backfill_falls_back_to_full_ttl_when_redis_ttl_missing( - self, handler_with_redis - ): + async def test_backfill_falls_back_to_full_ttl_when_redis_ttl_missing(self, handler_with_redis): cache_key = "{sensitive_route:hashed-key:session-123}:model" key = UserAPIKeyAuth(api_key="hashed-key") handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( return_value="on-premise-model" ) - handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = ( - AsyncMock(return_value=None) - ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = AsyncMock(return_value=None) await handler_with_redis._get_routed_model("session-123", key) - assert ( - handler_with_redis.internal_usage_cache._ttls[cache_key] - == handler_with_redis.ttl - ) + assert handler_with_redis.internal_usage_cache._ttls[cache_key] == handler_with_redis.ttl @pytest.mark.asyncio async def test_get_routed_model_redis_fallback_on_error(self, handler_with_redis): handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( side_effect=Exception("Redis connection error") ) - handler_with_redis.internal_usage_cache._cache[ - "{sensitive_route:hashed-key:session-123}:model" - ] = "fallback-model" - result = await handler_with_redis._get_routed_model( - "session-123", UserAPIKeyAuth(api_key="hashed-key") + handler_with_redis.internal_usage_cache._cache["{sensitive_route:hashed-key:session-123}:model"] = ( + "fallback-model" ) + result = await handler_with_redis._get_routed_model("session-123", UserAPIKeyAuth(api_key="hashed-key")) assert result == "fallback-model" @pytest.mark.asyncio async def test_set_session_routing_with_redis(self, handler_with_redis): - handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache = ( - AsyncMock() - ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache = AsyncMock() await handler_with_redis.set_session_routing( session_id="session-456", model="on-premise-model", @@ -642,9 +635,7 @@ class TestRedisCache: handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache.assert_called_once() @pytest.mark.asyncio - async def test_set_session_routing_redis_fallback_on_error( - self, handler_with_redis - ): + async def test_set_session_routing_redis_fallback_on_error(self, handler_with_redis): handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache = AsyncMock( side_effect=Exception("Redis connection error") ) @@ -654,10 +645,7 @@ class TestRedisCache: user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), ) cache_key = "{sensitive_route:hashed-key:session-789}:model" - assert ( - handler_with_redis.internal_usage_cache._cache[cache_key] - == "on-premise-model" - ) + assert handler_with_redis.internal_usage_cache._cache[cache_key] == "on-premise-model" class TestPreCallHookEdgeCases: @@ -746,16 +734,12 @@ class TestProxyHandleSensitiveDataRouteException: assert result["model"] == "on-premise-model" assert ( - await routing_hook._get_routed_model( - "sess-sticky", UserAPIKeyAuth(api_key="tenant-a") - ) + await routing_hook._get_routed_model("sess-sticky", UserAPIKeyAuth(api_key="tenant-a")) == "on-premise-model" ) @pytest.mark.asyncio - async def test_non_sticky_routing_does_not_persist_override( - self, proxy_logging, routing_hook - ): + async def test_non_sticky_routing_does_not_persist_override(self, proxy_logging, routing_hook): proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook exc = SensitiveDataRouteException( route_to_model="on-premise-model", @@ -770,17 +754,10 @@ class TestProxyHandleSensitiveDataRouteException: ) assert result["model"] == "on-premise-model" - assert ( - await routing_hook._get_routed_model( - "sess-non-sticky", UserAPIKeyAuth(api_key="tenant-a") - ) - is None - ) + assert await routing_hook._get_routed_model("sess-non-sticky", UserAPIKeyAuth(api_key="tenant-a")) is None @pytest.mark.asyncio - async def test_sticky_routing_handles_none_user_api_key_dict( - self, proxy_logging, routing_hook - ): + async def test_sticky_routing_handles_none_user_api_key_dict(self, proxy_logging, routing_hook): proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook exc = SensitiveDataRouteException( route_to_model="on-premise-model", @@ -790,20 +767,13 @@ class TestProxyHandleSensitiveDataRouteException: ) data = {"model": "gpt-4", "metadata": {"session_id": "sess-no-key"}} - result = await proxy_logging._handle_sensitive_data_route_exception( - exc, data, None - ) + result = await proxy_logging._handle_sensitive_data_route_exception(exc, data, None) assert result["model"] == "on-premise-model" - assert ( - await routing_hook._get_routed_model("sess-no-key", None) - == "on-premise-model" - ) + assert await routing_hook._get_routed_model("sess-no-key", None) == "on-premise-model" @pytest.mark.asyncio - async def test_sticky_routing_scopes_jwt_users_by_principal( - self, proxy_logging, routing_hook - ): + async def test_sticky_routing_scopes_jwt_users_by_principal(self, proxy_logging, routing_hook): proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook exc = SensitiveDataRouteException( route_to_model="on-premise-model", @@ -875,7 +845,6 @@ class _RecordingGuardrail(CustomGuardrail): async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): self.ran = True - return None class _BlockingGuardrail(CustomGuardrail): @@ -887,9 +856,7 @@ class _BlockingGuardrail(CustomGuardrail): from litellm.exceptions import GuardrailRaisedException self.ran = True - raise GuardrailRaisedException( - message="blocked", guardrail_name=self.guardrail_name - ) + raise GuardrailRaisedException(message="blocked", guardrail_name=self.guardrail_name) class TestPreCallHookDeferredRouting: @@ -976,9 +943,7 @@ class TestPreCallHookDeferredRouting: from litellm.types.services import ServiceTypes class _SlowRoutingGuardrail(CustomGuardrail): - async def async_pre_call_hook( - self, user_api_key_dict, cache, data, call_type - ): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): await asyncio.sleep(0.02) self.handle_sensitive_data_detection(request_data=data) @@ -1008,9 +973,7 @@ class TestPreCallHookDeferredRouting: assert recorded.call_args.kwargs["service"] == ServiceTypes.PROXY_PRE_CALL @pytest.mark.asyncio - async def test_routing_recorded_as_intervention_not_prometheus_error( - self, proxy_logging - ): + async def test_routing_recorded_as_intervention_not_prometheus_error(self, proxy_logging): import litellm from litellm.integrations.prometheus import PrometheusLogger diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 292b2904546..d3037e9ccfa 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -27,6 +27,7 @@ from fastapi.testclient import TestClient import litellm import litellm.proxy.proxy_server as proxy_server_module +from litellm._internal_context import current_service_target from litellm.caching.caching import RedisCache from litellm.caching.redis_cluster_cache import RedisClusterCache from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded @@ -15489,3 +15490,47 @@ async def test_spend_capture_rate_check_job_clears_the_gauge_once_the_setting_is call(api_provider="openai", capture_rate=0.97), call(api_provider="openai", capture_rate=None), ] + + +@pytest.mark.asyncio +async def test_update_cache_reads_and_writes_declare_the_auth_objects_key_family(): + """The post-call spend write-back reads and rewrites the cached auth objects, so its + Redis spans must read ``redis.mget auth_objects`` / ``redis.set auth_objects`` (the key + family the auth phase declares) and the global spend scalar ``redis.set spend_counters``, + never a bare ``redis.mget`` with no owner.""" + from litellm.caching.caching import DualCache + + original_cache = litellm.proxy.proxy_server.user_api_key_cache + cache = DualCache() + setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache) + seen: list[tuple[str, str | None]] = [] + + async def _mget(keys, **_kwargs): + seen.append(("mget", current_service_target())) + return [{"user_id": "u1", "spend": 1.0} for _ in keys] + + async def _set_pipeline(**_kwargs): + seen.append(("set", current_service_target())) + + try: + with ( + patch.object(cache, "async_batch_get_cache", new=AsyncMock(side_effect=_mget)), + patch.object(cache, "async_set_cache_pipeline", new=AsyncMock(side_effect=_set_pipeline)), + ): + await litellm.proxy.proxy_server.update_cache( + token=None, + user_id="u1", + end_user_id=None, + team_id=None, + response_cost=2.0, + parent_otel_span=None, + ) + pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()] + if pending: + await asyncio.wait(pending, timeout=5) + finally: + setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache) + + assert seen, "update_cache must touch the cache for a priced user request" + assert {target for _, target in seen} == {"auth_objects"} + assert current_service_target() is None diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py index 761835078f4..0d408de9ec6 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py @@ -19,6 +19,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest import litellm.proxy.utils as utils_mod +from litellm._internal_context import current_service_target from litellm.proxy.utils import ( _config_cache_key, _ConfigRow, @@ -265,3 +266,31 @@ async def test_prefetch_config_params_swallows_db_error_without_caching( prisma.db.litellm_config.find_many = AsyncMock(side_effect=RuntimeError("boom")) await prefetch_config_params(prisma, ["a", "b"]) assert _swap_config_cache._store == {} + + +@pytest.mark.asyncio +async def test_config_param_cache_calls_declare_the_config_params_key_family( + _swap_config_cache: Any, +) -> None: + """The config cache read and the miss write-back both run inside + ``service_target("config_params")`` so their Redis spans read + ``redis.get config_params`` / ``redis.set config_params``.""" + seen: list[tuple[str, Any]] = [] + + async def _get(*_args: Any, **_kwargs: Any) -> None: + seen.append(("get", current_service_target())) + + async def _set(*_args: Any, **_kwargs: Any) -> None: + seen.append(("set", current_service_target())) + + _swap_config_cache.async_get_cache = AsyncMock(side_effect=_get) + _swap_config_cache.async_set_cache = AsyncMock(side_effect=_set) + prisma = MagicMock() + prisma.get_generic_data = AsyncMock( + return_value=SimpleNamespace(param_name="p1", param_value={"x": 1}) + ) + + await get_config_param(prisma, "p1") + + assert seen == [("get", "config_params"), ("set", "config_params")] + assert current_service_target() is None diff --git a/tests/unit/router_utils/test_cooldown_cache.py b/tests/unit/router_utils/test_cooldown_cache.py index 6f90fa8465f..06dd294fc11 100644 --- a/tests/unit/router_utils/test_cooldown_cache.py +++ b/tests/unit/router_utils/test_cooldown_cache.py @@ -8,7 +8,7 @@ from unittest.mock import MagicMock import pytest # Add the parent directory to the system path - +from litellm._internal_context import current_service_target from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker @@ -582,3 +582,48 @@ class TestCooldownSurvivesUnrelatedCacheTraffic: assert [model_id] == [entry[0] for entry in active], ( "unrelated router cache traffic must not evict a cooldown that is still running" ) + + +class TestCooldownStoreCallsDeclareTheirKeyFamily: + """Every cooldown store call runs inside ``service_target("router_cooldowns")`` so the + Redis service spans read ``redis.set router_cooldowns`` / ``redis.mget router_cooldowns``, + the sync paths included (the async MGET already did).""" + + def _cooldown_cache_with_recording_store(self, seen: list[tuple[str, str | None]]) -> CooldownCache: + cc = CooldownCache(cache=DualCache(in_memory_cache=InMemoryCache()), default_cooldown_time=60.0) + store = MagicMock() + + def _set_cache(**_kwargs): + seen.append(("set", current_service_target())) + + def _batch_get_cache(**_kwargs): + seen.append(("mget", current_service_target())) + return [] + + store.set_cache.side_effect = _set_cache + store.batch_get_cache.side_effect = _batch_get_cache + cc._cooldown_store = store + return cc + + def test_sync_cooldown_write_runs_under_router_cooldowns(self): + seen: list[tuple[str, str | None]] = [] + cc = self._cooldown_cache_with_recording_store(seen) + + cc.add_deployment_to_cooldown( + model_id="dep-1", + original_exception=Exception("Internal server error"), + exception_status=500, + cooldown_time=30.0, + ) + + assert seen == [("set", "router_cooldowns")] + assert current_service_target() is None + + def test_sync_cooldown_reads_run_under_router_cooldowns(self): + seen: list[tuple[str, str | None]] = [] + cc = self._cooldown_cache_with_recording_store(seen) + + assert cc.get_active_cooldowns(["dep-1"], parent_otel_span=None) == [] + assert cc.get_min_cooldown(["dep-1"], parent_otel_span=None) == 60.0 + + assert seen == [("mget", "router_cooldowns"), ("mget", "router_cooldowns")] diff --git a/tests/unit/router_utils/test_cooldown_handlers.py b/tests/unit/router_utils/test_cooldown_handlers.py index 6fed4be5909..5526a38a646 100644 --- a/tests/unit/router_utils/test_cooldown_handlers.py +++ b/tests/unit/router_utils/test_cooldown_handlers.py @@ -1,10 +1,12 @@ from unittest.mock import MagicMock, patch import litellm +from litellm._internal_context import current_service_target from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.router_utils.cooldown_handlers import ( _get_deployment_cooldown_policy, + _increment_allowed_fails, _resolve_allowed_fails_from_policy, _should_cooldown_based_on_deployment_policy, should_cooldown_based_on_allowed_fails_policy, @@ -501,3 +503,35 @@ class TestTeamModelCooldownAlternatives: ) is False ) + + +class TestIncrementAllowedFailsServiceTarget: + def test_fail_counter_bump_declares_the_router_cooldowns_key_family(self): + """The allowed_fails INCR is cooldown bookkeeping, so its service span must read + ``redis.incr router_cooldowns`` rather than a bare ``redis.incr``.""" + seen: list[str | None] = [] + cache = MagicMock(spec=DualCache) + + def _increment(**_kwargs): + seen.append(current_service_target()) + return 2 + + cache.increment_cache.side_effect = _increment + + assert _increment_allowed_fails(cache, "deployment:dep-1:fails", ttl=60.0) == 2 + assert seen == ["router_cooldowns"] + assert current_service_target() is None + + def test_in_memory_fallback_reads_under_the_same_target(self): + seen: list[str | None] = [] + cache = MagicMock(spec=DualCache) + cache.increment_cache.side_effect = ConnectionError("redis down") + + def _get(**_kwargs): + seen.append(current_service_target()) + return 4 + + cache.get_cache.side_effect = _get + + assert _increment_allowed_fails(cache, "deployment:dep-1:fails", ttl=60.0) == 4 + assert seen == ["router_cooldowns"] diff --git a/tests/unit/test_internal_context.py b/tests/unit/test_internal_context.py new file mode 100644 index 00000000000..295d2e51023 --- /dev/null +++ b/tests/unit/test_internal_context.py @@ -0,0 +1,261 @@ +"""``with_service_target`` and ``service_caller`` carry the purpose and the caller of a datastore call +to code that cannot see them from its own frames, and every Redis producer on the proxy request path +declares a key family so no request-path span renders as a bare ``redis.get``.""" + +import ast +import asyncio +import contextvars +import re +from collections.abc import Generator +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import pytest + +from litellm._internal_context import ( + current_service_caller, + current_service_target, + service_caller, + service_target, + with_service_target, +) + +_REPO: Final = Path(__file__).resolve().parents[2] + +_REDIS_PRODUCER_ROOTS: Final = ("litellm", "enterprise") +# The cache implementations and facades: they emit the service events, their callers declare the family. +_CACHE_LAYER_DIRS: Final = ("litellm/caching", "litellm/_v2/cache") +# Helpers that act on a cache handed in by the declaring caller, or forward to the response-cache facade. +_CACHE_PARAMETER_HELPERS: Final = frozenset( + { + "litellm/proxy/common_utils/cache_coordinator.py", + "litellm/proxy/common_utils/user_api_key_cache.py", + "litellm/utils.py", + } +) +# Callers whose every cache call hits a process-local ``InMemoryCache`` (a ``DualCache`` built without +# ``redis_cache``, a ``local_only=True`` call, the client / logger / tool-name caches), so no Redis span exists. +_IN_MEMORY_ONLY_CALLERS: Final = frozenset( + { + "litellm/integrations/datadog/datadog_team_handler.py", + "litellm/integrations/humanloop.py", + "litellm/integrations/langfuse/langfuse_handler.py", + "litellm/integrations/langfuse/langfuse_prompt_management.py", + "litellm/integrations/newrelic/newrelic_team_handler.py", + "litellm/integrations/shadow_eval_logger.py", + "litellm/litellm_core_utils/litellm_logging.py", + "litellm/litellm_core_utils/prompt_templates/factory.py", + "litellm/litellm_core_utils/prompt_templates/image_handling.py", + "litellm/litellm_core_utils/specialty_caches/dynamic_logging_cache.py", + "litellm/litellm_core_utils/specialty_caches/service_trace_id_cache.py", + "litellm/llms/azure/common_utils.py", + "litellm/llms/bedrock/base_aws_llm.py", + "litellm/llms/custom_httpx/http_handler.py", + "litellm/llms/gigachat/authenticator.py", + "litellm/llms/litellm_proxy/skills/handler.py", + "litellm/llms/openai/common_utils.py", + "litellm/llms/openai_like/model_info.py", + "litellm/llms/vertex_ai/vertex_ai_non_gemini.py", + "litellm/llms/watsonx/common_utils.py", + "litellm/proxy/_experimental/mcp_server/byok_credential_cache.py", + "litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py", + "litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py", + "litellm/proxy/_experimental/mcp_server/operations.py", + "litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py", + "litellm/proxy/_experimental/mcp_server/outbound_credentials/token_endpoint.py", + "litellm/proxy/agent_endpoints/databricks_oauth.py", + "litellm/proxy/common_utils/registry_read_through.py", + "litellm/proxy/container_endpoints/ownership.py", + "litellm/proxy/discovery_endpoints/agent_skills_endpoints.py", + "litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py", + "litellm/proxy/spend_tracking/key_metadata_recovery.py", + "litellm/proxy/ui_crud_endpoints/latest_release_endpoints.py", + "litellm/responses/litellm_completion_transformation/transformation.py", + "litellm/router_utils/client_initalization_utils.py", + "litellm/router_utils/router_callbacks/track_deployment_metrics.py", + "litellm/secret_managers/cyberark_secret_manager.py", + "litellm/secret_managers/google_secret_manager.py", + "litellm/secret_managers/hashicorp_secret_manager.py", + "litellm/secret_managers/main.py", + } +) + +_CACHE_CALL: Final = re.compile( + r"\.(?:async_)?(?:get_cache|set_cache|batch_get_cache|batch_get_cache_shared|increment_cache|increment" + r"|set_cache_pipeline|set_cache_pipeline_with_ttls|set_cache_sadd|delete_cache|batch_set_cache|increment_pipeline" + r"|rpush|lpop|scan_iter|get_ttl|mget)\(" + r"|\b(?:reserve_redis_batch_reads|declare_batch_get|_prepare_batch_get)\(" + r"|\bbatch\.(?:set|delete|script|increment)\(" +) +_DECLARES_TARGET: Final = re.compile(r"\b(?:with_service_target|service_target|response_cache_phase)\(") +_BUILDS_A_REDIS_CACHE: Final = re.compile(r"\bRedisCache\(|\bredis_cache=(?!None\b)") + + +def _redis_producers() -> tuple[str, ...]: + files: Final = tuple( + path for root in _REDIS_PRODUCER_ROOTS for path in sorted((_REPO / root).rglob("*.py")) + ) # comprehension-ok: flatten the producer roots + relative: Final = tuple( + path.relative_to(_REPO).as_posix() for path in files if _CACHE_CALL.search(path.read_text()) + ) + return tuple(name for name in relative if not name.startswith(_CACHE_LAYER_DIRS)) + + +def test_every_redis_producer_declares_a_key_family() -> None: + """A module that reads or writes a shared cache without a declared target renders as a + bare ``redis.get`` / ``redis.mget`` (flat under the request span, or an unnamed INTERNAL root + for a background job), which is exactly what the sensitive-data pin read, the rate-limiter + MGET and the budget-reset job did in production. Only process-local callers are exempt.""" + exempt: Final = _CACHE_PARAMETER_HELPERS | _IN_MEMORY_ONLY_CALLERS + undeclared: Final = tuple( + name + for name in _redis_producers() + if name not in exempt and not _DECLARES_TARGET.search((_REPO / name).read_text()) + ) + assert undeclared == () + + +def test_every_in_memory_exemption_still_only_touches_a_process_local_cache() -> None: + """The exemption list is a claim about each file, so a file that is deleted or starts building + or receiving a ``RedisCache`` has to leave the list (and declare a family) rather than stay exempt.""" + producers: Final = frozenset(_redis_producers()) + stale: Final = tuple(sorted(_IN_MEMORY_ONLY_CALLERS - producers)) + assert stale == () + redis_backed: Final = tuple( + name for name in sorted(_IN_MEMORY_ONLY_CALLERS) if _BUILDS_A_REDIS_CACHE.search((_REPO / name).read_text()) + ) + assert redis_backed == () + + +def test_with_service_target_sets_the_target_for_sync_and_async_calls_and_restores_it() -> None: + @with_service_target("rate_limits") + def read() -> str | None: + return current_service_target() + + @with_service_target("rate_limits") + async def read_async() -> str | None: + await asyncio.sleep(0) + return current_service_target() + + assert read() == "rate_limits" + assert asyncio.run(read_async()) == "rate_limits" + assert current_service_target() is None + with service_target("auth_objects"): + assert read() == "rate_limits" + assert current_service_target() == "auth_objects" + + +def test_with_service_target_keeps_the_wrapped_signature_and_coroutine_ness() -> None: + import inspect + + @with_service_target("rate_limits") + async def hook(self: object, data: dict[str, str], call_type: str) -> None: + return None + + assert inspect.iscoroutinefunction(hook) + assert tuple(inspect.signature(hook).parameters) == ("self", "data", "call_type") + assert hook.__name__ == "hook" + + +def test_service_caller_is_inherited_by_a_task_spawned_inside_it_and_cleared_after() -> None: + async def spawned() -> str | None: + return current_service_caller() + + async def main() -> tuple[str | None, str | None]: + with service_caller("prefetch <- auth"): + task = asyncio.create_task(spawned()) + return await task, current_service_caller() + + assert asyncio.run(main()) == ("prefetch <- auth", None) + + +@pytest.mark.parametrize("value", [None, "x"]) +def test_service_caller_restores_the_outer_value(value: str | None) -> None: + with service_caller(value): + with service_caller("inner"): + assert current_service_caller() == "inner" + assert current_service_caller() == value + assert current_service_caller() is None + + +class _Suspend: + def __await__(self) -> Generator[None]: + yield + + +def test_a_targeted_coroutine_closed_from_another_context_does_not_raise() -> None: + @with_service_target("router_usage") + async def sync_forever() -> None: + await _Suspend() + + suspended: Final = sync_forever() + contextvars.copy_context().run(suspended.send, None) + contextvars.copy_context().run(suspended.close) + assert current_service_target() is None + + +_DIRECT_REDIS_CALL: Final = re.compile(r"\b_?redis_cache\.(?!async_register_script\b)(?:async_)?\w+\(") + + +@dataclass(frozen=True, slots=True) +class _FunctionScan: + name: str + reaches_redis_directly: bool + declares_a_family: bool + referenced_names: frozenset[str] + + +def _scan_function(source: str, fn: ast.FunctionDef | ast.AsyncFunctionDef) -> _FunctionScan: + body: Final = ast.get_source_segment(source, fn) or "" + decorators: Final = "\n".join(ast.get_source_segment(source, d) or "" for d in fn.decorator_list) + nodes: Final = tuple(ast.walk(fn)) + names: Final = frozenset(n.id for n in nodes if isinstance(n, ast.Name)) + attrs: Final = frozenset(n.attr for n in nodes if isinstance(n, ast.Attribute)) + return _FunctionScan( + name=fn.name, + reaches_redis_directly=bool(_DIRECT_REDIS_CALL.search(body)), + declares_a_family=bool(_DECLARES_TARGET.search(body + "\n" + decorators)), + referenced_names=(names | attrs) - {fn.name}, + ) + + +def _covered_by_callers(scans: tuple[_FunctionScan, ...], covered: frozenset[str]) -> frozenset[str]: + """Close ``covered`` over functions whose every in-file caller already declares a family.""" + callers: Final = { + scan.name: frozenset( + other.name for other in scans if other.name != scan.name and scan.name in other.referenced_names + ) + for scan in scans + } + grown: Final = covered | frozenset( + name for name, callers_of in callers.items() if callers_of and callers_of <= covered + ) + return grown if grown == covered else _covered_by_callers(scans, grown) + + +def _direct_redis_callers_without_a_family(name: str) -> tuple[str, ...]: + source: Final = (_REPO / name).read_text() + scans: Final = tuple( + _scan_function(source, node) + for node in ast.walk(ast.parse(source)) + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + ) + declared: Final = frozenset(scan.name for scan in scans if scan.declares_a_family) + covered: Final = _covered_by_callers(scans, declared) + return tuple(f"{name}::{scan.name}" for scan in scans if scan.reaches_redis_directly and scan.name not in covered) + + +def test_every_function_that_reaches_redis_directly_declares_its_family() -> None: + """A file-level declaration hides the producer that lacks one: the Claude Code session router + binding read sat in ``router.py`` beside dozens of declared families and still shipped as a bare + ``redis.get``. A function that bypasses the cache facades and calls ``redis_cache`` itself must + carry the family on itself, its decorator, or every one of its in-file callers.""" + exempt_files: Final = _CACHE_PARAMETER_HELPERS | _IN_MEMORY_ONLY_CALLERS + undeclared: Final = tuple( + function + for name in _redis_producers() + if name not in exempt_files + for function in _direct_redis_callers_without_a_family(name) + ) # comprehension-ok: flatten per-file findings + assert undeclared == () diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 96dddf15869..aeab488d270 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -19,10 +19,12 @@ import openai import pytest import respx from fastapi import HTTPException +from opentelemetry import trace import litellm from litellm import Router from litellm.caching.caching import DualCache +from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import _redis_circuit_breaker_guard from litellm.exceptions import GuardrailRaisedException, MidStreamFallbackError, ModifyResponseException from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper @@ -18882,3 +18884,110 @@ async def test_router_subclass_overriding_async_get_healthy_deployments_with_the response: Final = await router.acompletion(model="m", messages=[{"role": "user", "content": "x"}]) assert response.choices[0].message.content == "hi" + + +@pytest.mark.asyncio +async def test_failure_rpm_increment_declares_the_router_usage_key_family(): + """The RPM bump a failed call still earns is router usage bookkeeping, so its Redis span + reads ``redis.incr router_usage`` rather than a bare ``redis.incr``.""" + from unittest.mock import AsyncMock + + from litellm._internal_context import current_service_target + + router = Router( + model_list=[ + { + "model_name": "gpt-group", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake", "mock_response": "hi"}, + "model_info": {"id": "dep-1"}, + } + ] + ) + seen: list[str | None] = [] + + async def _increment(**_kwargs): + seen.append(current_service_target()) + + with patch.object(router.cache, "async_increment_cache", new=AsyncMock(side_effect=_increment)): + await router.async_deployment_callback_on_failure( + kwargs={ + "call_type": "acompletion", + "litellm_params": { + "metadata": {"deployment": "openai/gpt-4o", "model_group": "gpt-group"}, + "model_info": {"id": "dep-1"}, + }, + }, + completion_response=None, + start_time=None, + end_time=None, + ) + + assert seen == ["router_usage"] + assert current_service_target() is None + +class _SpanRecordingInMemoryCache(InMemoryCache): + """Records the live OTel span each read runs under, so the test sees what a Redis span would nest in.""" + + def __init__(self) -> None: + super().__init__() + self.active_span_names: list[str] = [] + + async def async_batch_get_cache(self, keys, **kwargs): + self.active_span_names.append(trace.get_current_span().name) + return await super().async_batch_get_cache(keys, **kwargs) + + async def async_get_cache(self, key, **kwargs): + self.active_span_names.append(trace.get_current_span().name) + return await super().async_get_cache(key, **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 + + +@pytest.mark.asyncio +async def test_deployment_selection_runs_inside_a_route_phase_named_after_the_model_group(v2_span_exporter): + """Picking a deployment opens ``route {model_group}`` (the requested group, not the deployment + it picks) under the server span, and the cooldown reads it issues run inside it, so their Redis + spans nest there instead of lying flat under the request.""" + from opentelemetry.sdk.trace import TracerProvider + + router = Router( + model_list=[ + { + "model_name": "gpt-group", + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "fake", "mock_response": "a"}, + "model_info": {"id": "dep-a"}, + }, + { + "model_name": "gpt-group", + "litellm_params": {"model": "openai/gpt-5.4", "api_key": "fake", "mock_response": "b"}, + "model_info": {"id": "dep-b"}, + }, + ] + ) + recording_cache = _SpanRecordingInMemoryCache() + router.cache.in_memory_cache = recording_cache + router.cooldown_cache.cooldown_store.in_memory_cache = recording_cache + + with TracerProvider().get_tracer("test").start_as_current_span("POST /v1/chat/completions") as server_span: + deployment = await router.async_get_available_deployment(model="gpt-group", request_kwargs={}) + + assert deployment["model_info"]["id"] in {"dep-a", "dep-b"} + (route_span,) = v2_span_exporter.get_finished_spans() + assert route_span.name == "route gpt-group" + assert route_span.parent is not None and route_span.parent.span_id == server_span.get_span_context().span_id + assert route_span.end_time is not None + assert recording_cache.active_span_names and set(recording_cache.active_span_names) == {"route gpt-group"} diff --git a/tests/unit/test_service_logger.py b/tests/unit/test_service_logger.py index de46403b64d..2ee04cc3b96 100644 --- a/tests/unit/test_service_logger.py +++ b/tests/unit/test_service_logger.py @@ -200,8 +200,8 @@ async def test_service_span_emitted_for_v2_logger_in_service_callback(monkeypatc parent.end() names = [s.name for s in exporter.get_finished_spans()] - # Span name is "{service} {call_type}" so repeated calls stay distinguishable. - assert "redis async_set_cache" in names + # Span name is "{service}.{verb}" (the method rides on db.operation.name) so repeated calls stay distinguishable. + assert "redis.set" in names @pytest.mark.asyncio @@ -277,3 +277,42 @@ async def test_service_failure_span_not_duplicated_for_string_and_instance( s for s in exporter.get_finished_spans() if s.name == "postgres get_user_object" ] assert len(db_spans) == 1 + + +@pytest.mark.asyncio +async def test_only_redis_service_spans_carry_the_ambient_key_family(monkeypatch): + """A key family set for a Redis read must not label the DB write-back that a + task spawned inside that context performs later.""" + from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + + from litellm._internal_context import service_target + from litellm.integrations.otel.logger import OpenTelemetryV2 + from litellm.integrations.otel.model.config import OpenTelemetryV2Config + from litellm.integrations.otel.model.semconv import LiteLLM + from litellm.integrations.otel.plumbing import providers + + cfg = OpenTelemetryV2Config(exporter="in_memory") + exporter = InMemorySpanExporter() + otel = OpenTelemetryV2(config=cfg, tracer_provider=providers.build_tracer_provider(cfg, exporter=exporter)) + monkeypatch.setattr(litellm, "service_callback", [otel]) + service_logger = ServiceLogging() + start = datetime(2026, 2, 13, 22, 35, 0) + end = datetime(2026, 2, 13, 22, 35, 1) + + with service_target("router_session_pins"): + await service_logger.async_service_success_hook( + service=ServiceTypes.REDIS, call_type="async_get_cache", duration=1.0, start_time=start, end_time=end + ) + await service_logger.async_service_success_hook( + service=ServiceTypes.BATCH_WRITE_TO_DB, + call_type="_PROXY_track_cost_callback", + duration=1.0, + start_time=start, + end_time=end, + ) + + targets = {span.name: span.attributes.get(LiteLLM.SERVICE_TARGET) for span in exporter.get_finished_spans()} + assert targets == { + "redis.get router_session_pins": "router_session_pins", + "batch_write_to_db _PROXY_track_cost_callback": None, + }