mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Perf speed up pytest (#15951)
* perf: Skip sleep delays in base_mail.py during tests to improve test speed * perf: Mock datetime.now in parallel_request_limiter_v3.py to improve test speed * pref: Mock urllib system calls in test_aiohttp_transport.py to improve test speed * chore: add --durations=50 to visualize slowest tests * pref: reduce setup phase overhead by widening fixture scope in conftest.py * test: stabilize flaky tests * fix: minor issue
This commit is contained in:
parent
5ad108bc9b
commit
8b33328cc1
11 changed files with 116 additions and 37 deletions
3
.github/workflows/test-litellm.yml
vendored
3
.github/workflows/test-litellm.yml
vendored
|
|
@ -33,6 +33,7 @@ jobs:
|
|||
poetry run pip install "google-genai==1.22.0"
|
||||
poetry run pip install "google-cloud-aiplatform>=1.38"
|
||||
poetry run pip install "fastapi-offline==1.7.3"
|
||||
poetry run pip install "python-multipart==0.0.18"
|
||||
- name: Setup litellm-enterprise as local package
|
||||
run: |
|
||||
cd enterprise
|
||||
|
|
@ -40,4 +41,4 @@ jobs:
|
|||
cd ..
|
||||
- name: Run tests
|
||||
run: |
|
||||
poetry run pytest tests/test_litellm --tb=short -vv --maxfail=10 -n 4
|
||||
poetry run pytest tests/test_litellm --tb=short -vv --maxfail=10 -n 4 --durations=50
|
||||
|
|
|
|||
|
|
@ -3,7 +3,8 @@ Dynamic rate limiter v3 - Saturation-aware priority-based rate limiting
|
|||
"""
|
||||
|
||||
import os
|
||||
from typing import Dict, List, Literal, Optional, Union
|
||||
from datetime import datetime
|
||||
from typing import Callable, Dict, List, Literal, Optional, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -42,9 +43,15 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
- When saturated: strict priority-based limits enforced (fair)
|
||||
- Uses v3 limiter's atomic Lua scripts for race-free increments
|
||||
"""
|
||||
def __init__(self, internal_usage_cache: DualCache):
|
||||
def __init__(
|
||||
self,
|
||||
internal_usage_cache: DualCache,
|
||||
time_provider: Optional[Callable[[], datetime]] = None,
|
||||
):
|
||||
self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache)
|
||||
self.v3_limiter = _PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache)
|
||||
self.v3_limiter = _PROXY_MaxParallelRequestsHandler_v3(
|
||||
self.internal_usage_cache, time_provider=time_provider
|
||||
)
|
||||
|
||||
def update_variables(self, llm_router: Router):
|
||||
self.llm_router = llm_router
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from math import floor
|
|||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
|
|
@ -137,8 +138,13 @@ class RateLimitResponseWithDescriptors(TypedDict):
|
|||
|
||||
|
||||
class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
||||
def __init__(self, internal_usage_cache: InternalUsageCache):
|
||||
def __init__(
|
||||
self,
|
||||
internal_usage_cache: InternalUsageCache,
|
||||
time_provider: Optional[Callable[[], datetime]] = None,
|
||||
):
|
||||
self.internal_usage_cache = internal_usage_cache
|
||||
self._time_provider = time_provider or datetime.now
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
self.batch_rate_limiter_script = (
|
||||
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
|
||||
|
|
@ -156,6 +162,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
|
||||
|
||||
def _get_current_time(self) -> datetime:
|
||||
"""Return the current time for rate limiting calculations."""
|
||||
return self._time_provider()
|
||||
|
||||
def _is_redis_cluster(self) -> bool:
|
||||
"""
|
||||
Check if the dual cache is using Redis cluster.
|
||||
|
|
@ -425,7 +435,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
read_only: If True, only check limits without incrementing counters
|
||||
"""
|
||||
|
||||
now = datetime.now().timestamp()
|
||||
current_time = self._get_current_time()
|
||||
now = current_time.timestamp()
|
||||
now_int = int(now) # Convert to integer for Redis Lua script
|
||||
|
||||
# Collect all keys and their metadata upfront
|
||||
|
|
@ -1090,7 +1101,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
descriptor = descriptors[floor(i / 2)]
|
||||
|
||||
# Calculate reset time (window_start + window_size)
|
||||
now = datetime.now().timestamp()
|
||||
now = self._get_current_time().timestamp()
|
||||
reset_time = now + self.window_size # Conservative estimate
|
||||
reset_time_formatted = datetime.fromtimestamp(
|
||||
reset_time
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ def event_loop():
|
|||
|
||||
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
|
||||
|
|
|
|||
|
|
@ -4,13 +4,11 @@ import sys
|
|||
import unittest.mock as mock
|
||||
from unittest.mock import patch
|
||||
|
||||
from enterprise.litellm_enterprise.enterprise_callbacks.send_emails.base_email import BaseEmailLogger
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import (
|
||||
BaseEmailLogger,
|
||||
)
|
||||
from litellm_enterprise.types.enterprise_callbacks.send_emails import (
|
||||
EmailEvent,
|
||||
SendKeyCreatedEmailEvent,
|
||||
|
|
@ -20,6 +18,13 @@ from litellm.integrations.email_templates.email_footer import EMAIL_FOOTER
|
|||
from litellm.proxy._types import Litellm_EntityType, WebhookEvent
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def no_invitation_wait(monkeypatch):
|
||||
async def _noop(self):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(BaseEmailLogger, "_wait_for_invitation_creation", _noop)
|
||||
|
||||
@pytest.fixture
|
||||
def base_email_logger():
|
||||
return BaseEmailLogger()
|
||||
|
|
|
|||
|
|
@ -188,6 +188,8 @@ async def test_handle_async_request_uses_env_proxy(monkeypatch):
|
|||
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
|
||||
monkeypatch.setenv("https_proxy", proxy_url)
|
||||
monkeypatch.delenv("DISABLE_AIOHTTP_TRUST_ENV", raising=False)
|
||||
monkeypatch.setattr("urllib.request.getproxies", lambda: {"http": proxy_url, "https": proxy_url})
|
||||
monkeypatch.setattr("urllib.request.proxy_bypass", lambda host: False)
|
||||
|
||||
captured = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -11,6 +11,13 @@ import pytest
|
|||
|
||||
import litellm
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_litellm_http_client_cache():
|
||||
"""Ensure each test gets a fresh async HTTP client mock."""
|
||||
from litellm import in_memory_llm_clients_cache
|
||||
|
||||
in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
class TestVertexGemmaCompletion:
|
||||
"""Test completion flow for Vertex AI Gemma models using litellm.acompletion()"""
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import asyncio
|
|||
import os
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -22,6 +23,24 @@ from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
|
|||
)
|
||||
|
||||
|
||||
class TimeController:
|
||||
def __init__(self):
|
||||
self._current = datetime.utcnow()
|
||||
|
||||
def now(self) -> datetime:
|
||||
return self._current
|
||||
|
||||
def advance(self, seconds: float) -> None:
|
||||
self._current += timedelta(seconds=seconds)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def time_controller(monkeypatch):
|
||||
controller = TimeController()
|
||||
monkeypatch.setattr(time, "time", lambda: controller.now().timestamp())
|
||||
return controller
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_priority_weight_allocation():
|
||||
"""
|
||||
|
|
@ -195,7 +214,7 @@ async def test_concurrent_priority_requests():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_100_concurrent_priority_requests():
|
||||
async def test_100_concurrent_priority_requests(time_controller):
|
||||
"""
|
||||
Stress test: 100 concurrent requests with mixed priorities over 10 seconds.
|
||||
|
||||
|
|
@ -211,7 +230,9 @@ async def test_100_concurrent_priority_requests():
|
|||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
||||
dual_cache = DualCache()
|
||||
handler = DynamicRateLimitHandler(internal_usage_cache=dual_cache)
|
||||
handler = DynamicRateLimitHandler(
|
||||
internal_usage_cache=dual_cache, time_provider=time_controller.now
|
||||
)
|
||||
|
||||
model = "stress-test-model"
|
||||
total_tpm = 1000
|
||||
|
|
@ -307,7 +328,8 @@ async def test_100_concurrent_priority_requests():
|
|||
|
||||
# Add small delay between batches to spread over ~10 seconds
|
||||
if batch_idx < len(batches) - 1: # Don't sleep after last batch
|
||||
await asyncio.sleep(1.0) # 1 second between batches
|
||||
await asyncio.sleep(0)
|
||||
time_controller.advance(1.0) # simulate 1s passing between batches
|
||||
|
||||
end_time = time.time()
|
||||
total_duration = end_time - start_time
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ Unit Tests for the max parallel request limiter v3 for the proxy
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import pytest
|
||||
|
|
@ -21,10 +22,27 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
|||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
class TimeController:
|
||||
def __init__(self):
|
||||
self._current = datetime.utcnow()
|
||||
|
||||
def now(self) -> datetime:
|
||||
return self._current
|
||||
|
||||
def advance(self, seconds: float) -> None:
|
||||
self._current += timedelta(seconds=seconds)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def time_controller(monkeypatch):
|
||||
controller = TimeController()
|
||||
monkeypatch.setattr(time, "time", lambda: controller.now().timestamp())
|
||||
return controller
|
||||
|
||||
|
||||
@pytest.mark.flaky(reruns=3)
|
||||
@pytest.mark.asyncio
|
||||
async def test_sliding_window_rate_limit_v3(monkeypatch):
|
||||
async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller):
|
||||
"""
|
||||
Test the sliding window rate limiting functionality
|
||||
"""
|
||||
|
|
@ -34,7 +52,8 @@ async def test_sliding_window_rate_limit_v3(monkeypatch):
|
|||
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, rpm_limit=3)
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
internal_usage_cache=InternalUsageCache(local_cache),
|
||||
time_provider=time_controller.now,
|
||||
)
|
||||
|
||||
# Mock the batch_rate_limiter_script to simulate window expiry and use correct key construction
|
||||
|
|
@ -103,7 +122,7 @@ async def test_sliding_window_rate_limit_v3(monkeypatch):
|
|||
assert "Rate limit exceeded" in str(exc_info.value.detail)
|
||||
|
||||
# Wait for window to expire (2 seconds)
|
||||
await asyncio.sleep(3)
|
||||
time_controller.advance(3)
|
||||
|
||||
print("WAITED 3 seconds")
|
||||
|
||||
|
|
@ -116,7 +135,7 @@ async def test_sliding_window_rate_limit_v3(monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rate_limiter_script_return_values_v3(monkeypatch):
|
||||
async def test_rate_limiter_script_return_values_v3(monkeypatch, time_controller):
|
||||
"""
|
||||
Test that the rate limiter script returns both counter and window values correctly
|
||||
"""
|
||||
|
|
@ -126,7 +145,8 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch):
|
|||
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, rpm_limit=3)
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
internal_usage_cache=InternalUsageCache(local_cache),
|
||||
time_provider=time_controller.now,
|
||||
)
|
||||
|
||||
# Mock the batch_rate_limiter_script to simulate window expiry and use correct key construction
|
||||
|
|
@ -199,7 +219,7 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch):
|
|||
assert new_counter_value == 2, "Counter should be 2 after second request"
|
||||
|
||||
# Wait for window to expire
|
||||
await asyncio.sleep(3)
|
||||
time_controller.advance(3)
|
||||
|
||||
# Make request after window expiry
|
||||
await parallel_request_handler.async_pre_call_hook(
|
||||
|
|
@ -226,7 +246,7 @@ async def test_rate_limiter_script_return_values_v3(monkeypatch):
|
|||
)
|
||||
@pytest.mark.flaky(reruns=3)
|
||||
@pytest.mark.asyncio
|
||||
async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object):
|
||||
async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object, time_controller):
|
||||
"""
|
||||
Test normal router call with parallel request limiter v3 for TPM rate limiting
|
||||
"""
|
||||
|
|
@ -276,7 +296,8 @@ async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object):
|
|||
)
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
internal_usage_cache=InternalUsageCache(local_cache),
|
||||
time_provider=time_controller.now,
|
||||
)
|
||||
|
||||
# Mock the batch_rate_limiter_script to simulate window expiry and use correct key construction
|
||||
|
|
@ -359,7 +380,8 @@ async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object):
|
|||
},
|
||||
mock_response="hello",
|
||||
)
|
||||
await asyncio.sleep(1) # success is done in a separate thread
|
||||
await asyncio.sleep(0)
|
||||
time_controller.advance(1)
|
||||
|
||||
# Verify the token count is tracked
|
||||
counter_value = await local_cache.async_get_cache(key=counter_key)
|
||||
|
|
@ -383,7 +405,7 @@ async def test_normal_router_call_tpm_v3(monkeypatch, rate_limit_object):
|
|||
)
|
||||
|
||||
# Wait for window to expire
|
||||
await asyncio.sleep(3)
|
||||
time_controller.advance(3)
|
||||
|
||||
# Make request after window expiry
|
||||
await parallel_request_handler.async_pre_call_hook(
|
||||
|
|
|
|||
|
|
@ -125,16 +125,14 @@ def test_get_microsoft_callback_response():
|
|||
"surname": "User",
|
||||
}
|
||||
|
||||
future = asyncio.Future()
|
||||
future.set_result(mock_response)
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"MICROSOFT_CLIENT_SECRET": "mock_secret", "MICROSOFT_TENANT": "mock_tenant"},
|
||||
):
|
||||
mock_verify = AsyncMock(return_value=mock_response)
|
||||
with patch(
|
||||
"fastapi_sso.sso.microsoft.MicrosoftSSO.verify_and_process",
|
||||
return_value=future,
|
||||
new=mock_verify,
|
||||
):
|
||||
# Act
|
||||
result = asyncio.run(
|
||||
|
|
@ -166,15 +164,14 @@ def test_get_microsoft_callback_response_raw_sso_response():
|
|||
"surname": "User",
|
||||
}
|
||||
|
||||
future = asyncio.Future()
|
||||
future.set_result(mock_response)
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"MICROSOFT_CLIENT_SECRET": "mock_secret", "MICROSOFT_TENANT": "mock_tenant"},
|
||||
):
|
||||
mock_verify = AsyncMock(return_value=mock_response)
|
||||
with patch(
|
||||
"fastapi_sso.sso.microsoft.MicrosoftSSO.verify_and_process",
|
||||
return_value=future,
|
||||
new=mock_verify,
|
||||
):
|
||||
# Act
|
||||
result = asyncio.run(
|
||||
|
|
@ -207,12 +204,10 @@ def test_get_google_callback_response():
|
|||
"family_name": "User",
|
||||
}
|
||||
|
||||
future = asyncio.Future()
|
||||
future.set_result(mock_response)
|
||||
|
||||
with patch.dict(os.environ, {"GOOGLE_CLIENT_SECRET": "mock_secret"}):
|
||||
mock_verify = AsyncMock(return_value=mock_response)
|
||||
with patch(
|
||||
"fastapi_sso.sso.google.GoogleSSO.verify_and_process", return_value=future
|
||||
"fastapi_sso.sso.google.GoogleSSO.verify_and_process", new=mock_verify
|
||||
):
|
||||
# Act
|
||||
result = asyncio.run(
|
||||
|
|
@ -2072,4 +2067,3 @@ class TestPKCEFunctionality:
|
|||
assert "code_challenge=" in updated_location
|
||||
assert "code_challenge_method=S256" in updated_location
|
||||
assert f"state={test_state}" in updated_location
|
||||
|
||||
|
|
|
|||
|
|
@ -71,6 +71,14 @@ def disable_budget_sync(monkeypatch):
|
|||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_router_callbacks():
|
||||
"""Ensure router budget callbacks from previous tests do not leak state."""
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
yield
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ui_view_spend_logs_with_user_id(client, monkeypatch):
|
||||
# Mock data for the test
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue