From 4bb10ed0fc2ca5d42f4ce8dd77dac7b2fee9e18c Mon Sep 17 00:00:00 2001 From: mateo Date: Fri, 25 Sep 2026 00:10:23 +0000 Subject: [PATCH] fix(spend): use the cli-session alias for websearch spend, prometheus failure labels and the parallel limiter Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/prometheus.py | 5 ++- .../websearch_interception/handler.py | 2 +- .../proxy/hooks/parallel_request_limiter.py | 5 +-- .../integrations/test_prometheus_labels.py | 23 +++++++++++++ .../test_websearch_interception_handler.py | 33 +++++++++++++++++++ .../hooks/test_parallel_request_limiter.py | 23 +++++++++++++ 6 files changed, 87 insertions(+), 4 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index fb010ab5886..995f0683136 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -2605,6 +2605,7 @@ class PrometheusLogger(CustomLogger): from litellm.litellm_core_utils.litellm_logging import ( StandardLoggingPayloadSetup, ) + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup status_code: Final = self._extract_status_code(exception=original_exception) @@ -2623,7 +2624,9 @@ class PrometheusLogger(CustomLogger): end_user=user_api_key_dict.end_user_id, user=user_api_key_dict.user_id, user_email=user_api_key_dict.user_email, - hashed_api_key=None if status_code == 401 else user_api_key_dict.api_key, + hashed_api_key=None + if status_code == 401 + else LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict), api_key_alias=user_api_key_dict.key_alias, team=user_api_key_dict.team_id, team_alias=user_api_key_dict.team_alias, diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 90bbd5a00d8..6ebf485d717 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1687,7 +1687,7 @@ class WebSearchInterceptionLogger(CustomLogger): **user_api_key_metadata, **parent_correlation.as_search_metadata(), "model_group": search_tool_name, - "user_api_key": user_api_key_auth.api_key, + "user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_auth), "user_api_key_auth": user_api_key_auth, } diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index d41acadc4dd..e3485ebf25d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -20,6 +20,7 @@ from litellm.proxy.auth.auth_utils import ( from litellm.proxy.auth.budget_throttle import throttled_limit from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.utils import Usage if TYPE_CHECKING: @@ -250,7 +251,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): call_type: str, ): self.print_verbose("Inside Max Parallel Request Pre-Call Hook") - api_key: Final = user_api_key_dict.api_key + api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) max_parallel_requests = user_api_key_dict.max_parallel_requests if max_parallel_requests is None: max_parallel_requests = sys.maxsize @@ -803,7 +804,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): """ Retrieve the key's remaining rate limits. """ - api_key: Final = user_api_key_dict.api_key + api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) current_date: Final = datetime.now().strftime("%Y-%m-%d") current_hour: Final = datetime.now().strftime("%H") current_minute: Final = datetime.now().strftime("%M") diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/test_litellm/integrations/test_prometheus_labels.py index 41d0c44ff89..d598d59e183 100644 --- a/tests/test_litellm/integrations/test_prometheus_labels.py +++ b/tests/test_litellm/integrations/test_prometheus_labels.py @@ -716,6 +716,29 @@ async def test_failure_hook_emits_api_provider_value_on_failed_requests_metric() _clear_prometheus_registry() +@pytest.mark.asyncio +async def test_failure_hook_labels_a_cli_session_with_the_per_user_alias_not_the_login_token(): + from litellm.integrations.prometheus import PrometheusLogger + from litellm.proxy._types import UserAPIKeyAuth + + _clear_prometheus_registry() + try: + await PrometheusLogger().async_post_call_failure_hook( + request_data={"model": "gpt-4o-mini", "metadata": {}}, + original_exception=Exception("boom"), + user_api_key_dict=UserAPIKeyAuth( + api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa", + user_id="alice", + key_alias="cli-session-alice", + is_session_token=True, + ), + ) + hashed_keys = {s.labels.get("hashed_api_key") for s in _collected_samples("litellm_proxy_failed_requests_metric_total")} + assert hashed_keys == {"cli-session-alice"}, hashed_keys + finally: + _clear_prometheus_registry() + + async def _failed_requests_api_provider_labels( request_data: dict[str, object], original_exception: Exception, diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py index d6450d0f1de..a5ab28ba72a 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py @@ -230,6 +230,39 @@ async def test_execute_search_passes_selected_search_tool_litellm_params(monkeyp assert forwarded_kwargs["max_retries"] == 2 +@pytest.mark.asyncio +async def test_execute_search_attributes_cli_session_spend_to_the_per_user_alias_not_the_login_token(monkeypatch): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="perplexity-sonar-pro") + router = MagicMock() + router.search_tools = [ + { + "search_tool_name": "perplexity-sonar-pro", + "litellm_params": {"search_provider": "perplexity", "api_key": "fake-key"}, + } + ] + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + session = UserAPIKeyAuth( + api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa", + user_id="alice", + key_alias="cli-session-alice", + is_session_token=True, + ) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + + await logger._execute_search( + "what is litellm", + kwargs={"litellm_params": {"metadata": {"user_api_key_auth": session}}}, + ) + + forwarded_metadata = mock_asearch.await_args.kwargs["litellm_metadata"] + assert forwarded_metadata["user_api_key"] == "cli-session-alice" + assert forwarded_metadata["user_api_key_hash"] == "cli-session-alice" + + @pytest.mark.asyncio async def test_execute_search_attributes_spend_to_the_calling_key(monkeypatch): """An intercepted search is billed and logged against the key that made the LLM request. diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index 0e2683dcbfd..a83ddc69863 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -7,6 +7,7 @@ from datetime import datetime import pytest from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter import ( _PROXY_MaxParallelRequestsHandler, ) @@ -14,6 +15,28 @@ from litellm.proxy.utils import InternalUsageCache, hash_token from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage +@pytest.mark.asyncio +async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token(): + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache())) + session = UserAPIKeyAuth( + api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa", + user_id="alice", + key_alias="cli-session-alice", + is_session_token=True, + max_parallel_requests=5, + ) + + await handler.async_pre_call_hook( + user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion" + ) + + precise_minute = datetime.now().strftime("%Y-%m-%d-%H-%M") + counted = await handler.internal_usage_cache.async_get_cache( + key=f"cli-session-alice::{precise_minute}::request_count", litellm_parent_otel_span=None + ) + assert counted == {"current_requests": 1, "current_tpm": 0, "current_rpm": 1}, counted + + @pytest.mark.parametrize( "response_obj", [