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>
This commit is contained in:
mateo 2026-09-25 00:10:23 +00:00
parent 0ac9596ca6
commit 4bb10ed0fc
6 changed files with 87 additions and 4 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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