litellm/tests/unit/proxy/hooks/test_parallel_request_limiter.py
devin-ai-integration[bot] 39e31958f8
test(proxy): move auth, hooks, policy_engine and client tests into tests/unit/proxy (#43998)
* test(proxy): move auth, hooks, policy_engine and client tests into tests/unit/proxy

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): stub HIBP through respx by disabling the aiohttp transport

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): share the httpx transport fixture across proxy unit tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): restore proxy globals without a missing-value sentinel

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): package moved dirs and stub the login breach check at the HTTP boundary

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): isolate the mcp server manager per test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yuneng <yuneng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-10-01 10:11:45 -07:00

109 lines
3.7 KiB
Python

"""
Unit Tests for the max parallel request limiter v1 for the proxy
"""
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,
)
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",
[
EmbeddingResponse(
model="text-embedding-3-small",
usage=Usage(prompt_tokens=50, completion_tokens=0, total_tokens=50),
),
TextCompletionResponse(
model="gpt-3.5-turbo-instruct",
usage=Usage(prompt_tokens=20, completion_tokens=30, total_tokens=50),
),
],
)
@pytest.mark.asyncio
async def test_async_log_success_event_counts_non_chat_response_tokens(response_obj):
"""
Embedding and text completion responses must increment the per key, user,
team, and end user TPM counters, not just chat completion ModelResponse
objects.
"""
_api_key = hash_token("sk-12345")
user_id = "ishaan"
team_id = "litellm-team"
end_user_id = "customer-1"
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
scope_ids = [_api_key, user_id, team_id, end_user_id]
for scope_id in scope_ids:
await parallel_request_handler.internal_usage_cache.async_set_cache(
key=f"{scope_id}::{precise_minute}::request_count",
value={"current_requests": 1, "current_tpm": 0, "current_rpm": 1},
litellm_parent_otel_span=None,
)
kwargs = {
"litellm_params": {
"metadata": {
"user_api_key": _api_key,
"user_api_key_user_id": user_id,
"user_api_key_team_id": team_id,
"user_api_key_model_max_budget": {},
}
},
"user": end_user_id,
}
await parallel_request_handler.async_log_success_event(
kwargs=kwargs,
response_obj=response_obj,
start_time=datetime.now(),
end_time=datetime.now(),
)
for scope_id in scope_ids:
current = await parallel_request_handler.internal_usage_cache.async_get_cache(
key=f"{scope_id}::{precise_minute}::request_count",
litellm_parent_otel_span=None,
)
assert current["current_tpm"] == 50, (
f"expected 50 tokens counted for {scope_id}, "
f"got {current['current_tpm']}"
)