mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
0ac9596ca6
commit
4bb10ed0fc
6 changed files with 87 additions and 4 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue