mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Add queue and redis caps; pass gunicorn keepalive
Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
fd58c8c060
commit
78c2a39f1c
7 changed files with 215 additions and 16 deletions
|
|
@ -462,16 +462,14 @@ def get_redis_async_client(
|
|||
def get_redis_connection_pool(**env_overrides):
|
||||
redis_kwargs = _get_redis_client_logic(**env_overrides)
|
||||
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
|
||||
_normalize_max_connections(redis_kwargs=redis_kwargs)
|
||||
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
|
||||
pool_kwargs = {"timeout": REDIS_CONNECTION_POOL_TIMEOUT, "url": redis_kwargs["url"]}
|
||||
pool_kwargs = {
|
||||
"timeout": REDIS_CONNECTION_POOL_TIMEOUT,
|
||||
"url": redis_kwargs["url"],
|
||||
}
|
||||
if "max_connections" in redis_kwargs:
|
||||
try:
|
||||
pool_kwargs["max_connections"] = int(redis_kwargs["max_connections"])
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"REDIS: invalid max_connections value %r, ignoring",
|
||||
redis_kwargs["max_connections"],
|
||||
)
|
||||
pool_kwargs["max_connections"] = redis_kwargs["max_connections"]
|
||||
return async_redis.BlockingConnectionPool.from_url(**pool_kwargs)
|
||||
connection_class = async_redis.Connection
|
||||
if "ssl" in redis_kwargs:
|
||||
|
|
@ -483,6 +481,53 @@ def get_redis_connection_pool(**env_overrides):
|
|||
timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs
|
||||
)
|
||||
|
||||
|
||||
def _normalize_max_connections(redis_kwargs: dict) -> None:
|
||||
"""
|
||||
Normalize and clamp max_connections to avoid unbounded connection pools.
|
||||
|
||||
- Invalid values are ignored.
|
||||
- Values <= 0 are ignored.
|
||||
- Very large values are clamped by REDIS_MAX_CONNECTIONS_SOFT_CAP.
|
||||
"""
|
||||
if "max_connections" not in redis_kwargs:
|
||||
return
|
||||
|
||||
raw_max_connections = redis_kwargs.get("max_connections")
|
||||
try:
|
||||
max_connections = int(raw_max_connections)
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"REDIS: invalid max_connections value %r, ignoring",
|
||||
raw_max_connections,
|
||||
)
|
||||
redis_kwargs.pop("max_connections", None)
|
||||
return
|
||||
|
||||
if max_connections <= 0:
|
||||
verbose_logger.warning(
|
||||
"REDIS: max_connections must be > 0, got %r. Ignoring value.",
|
||||
raw_max_connections,
|
||||
)
|
||||
redis_kwargs.pop("max_connections", None)
|
||||
return
|
||||
|
||||
raw_soft_cap = get_secret("REDIS_MAX_CONNECTIONS_SOFT_CAP", default_value=10000) # type: ignore
|
||||
try:
|
||||
soft_cap = int(raw_soft_cap)
|
||||
except (TypeError, ValueError):
|
||||
soft_cap = 10000
|
||||
|
||||
if soft_cap > 0 and max_connections > soft_cap:
|
||||
verbose_logger.warning(
|
||||
"REDIS: max_connections=%s exceeds REDIS_MAX_CONNECTIONS_SOFT_CAP=%s. Clamping.",
|
||||
max_connections,
|
||||
soft_cap,
|
||||
)
|
||||
max_connections = soft_cap
|
||||
|
||||
redis_kwargs["max_connections"] = max_connections
|
||||
|
||||
def _pretty_print_redis_config(redis_kwargs: dict) -> None:
|
||||
"""Pretty print the Redis configuration using rich with sensitive data masking"""
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1334,6 +1334,9 @@ SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500))
|
|||
SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000))
|
||||
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
|
||||
SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
|
||||
# Hard cap for in-memory spend log buffering to prevent unbounded growth
|
||||
# when DB connectivity is degraded or unavailable.
|
||||
MAX_SPEND_LOG_QUEUE_SIZE = int(os.getenv("MAX_SPEND_LOG_QUEUE_SIZE", 10000))
|
||||
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS = int(
|
||||
os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)
|
||||
) # 1 minute
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from typing import (
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache, RedisCache
|
||||
from litellm.constants import DB_SPEND_UPDATE_JOB_NAME
|
||||
from litellm.constants import DB_SPEND_UPDATE_JOB_NAME, MAX_SPEND_LOG_QUEUE_SIZE
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy._types import (
|
||||
DB_CONNECTION_ERROR_TYPES,
|
||||
|
|
@ -678,16 +678,35 @@ class DBSpendUpdateWriter:
|
|||
payload.get("request_id"), payload.get("spend")
|
||||
)
|
||||
)
|
||||
if prisma_client is not None and spend_logs_url is not None:
|
||||
async with prisma_client._spend_log_transactions_lock:
|
||||
prisma_client.spend_log_transactions.append(payload)
|
||||
elif prisma_client is not None:
|
||||
async with prisma_client._spend_log_transactions_lock:
|
||||
prisma_client.spend_log_transactions.append(payload)
|
||||
else:
|
||||
if prisma_client is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"prisma_client is None. Skipping writing spend logs to db."
|
||||
)
|
||||
return prisma_client
|
||||
|
||||
# Keep spend_logs_url behavior unchanged; both paths buffer in memory first.
|
||||
_ = spend_logs_url
|
||||
async with prisma_client._spend_log_transactions_lock:
|
||||
queue_size = len(prisma_client.spend_log_transactions)
|
||||
if queue_size >= MAX_SPEND_LOG_QUEUE_SIZE:
|
||||
now = time.time()
|
||||
last_warn_ts = float(
|
||||
getattr(prisma_client, "_last_spend_log_queue_drop_warning_ts", 0.0)
|
||||
)
|
||||
# Throttle warning logs to avoid log storms under sustained DB outages.
|
||||
if now - last_warn_ts >= 30:
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking queue is full (%s items). "
|
||||
"Dropping new spend log entry. "
|
||||
"Set MAX_SPEND_LOG_QUEUE_SIZE to tune this limit.",
|
||||
MAX_SPEND_LOG_QUEUE_SIZE,
|
||||
)
|
||||
setattr(
|
||||
prisma_client, "_last_spend_log_queue_drop_warning_ts", now
|
||||
)
|
||||
return prisma_client
|
||||
|
||||
prisma_client.spend_log_transactions.append(payload)
|
||||
|
||||
return prisma_client
|
||||
|
||||
|
|
|
|||
|
|
@ -192,6 +192,7 @@ class ProxyInitializationHelpers:
|
|||
num_workers: int,
|
||||
ssl_certfile_path: str,
|
||||
ssl_keyfile_path: str,
|
||||
keepalive_timeout: Optional[int] = None,
|
||||
max_requests_before_restart: Optional[int] = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -276,6 +277,8 @@ class ProxyInitializationHelpers:
|
|||
# Optional: recycle workers after N requests to mitigate memory growth
|
||||
if max_requests_before_restart is not None:
|
||||
gunicorn_options["max_requests"] = max_requests_before_restart
|
||||
if keepalive_timeout is not None:
|
||||
gunicorn_options["keepalive"] = keepalive_timeout
|
||||
|
||||
if ssl_certfile_path is not None and ssl_keyfile_path is not None:
|
||||
print( # noqa
|
||||
|
|
@ -911,6 +914,7 @@ def run_server( # noqa: PLR0915
|
|||
num_workers=num_workers,
|
||||
ssl_certfile_path=ssl_certfile_path,
|
||||
ssl_keyfile_path=ssl_keyfile_path,
|
||||
keepalive_timeout=keepalive_timeout,
|
||||
max_requests_before_restart=max_requests_before_restart,
|
||||
)
|
||||
elif run_hypercorn is True:
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ class MockPrismaClient:
|
|||
# Add lock for spend_log_transactions (matches real PrismaClient)
|
||||
import asyncio
|
||||
self._spend_log_transactions_lock = asyncio.Lock()
|
||||
self._last_spend_log_queue_drop_warning_ts = 0.0
|
||||
|
||||
def jsonify_object(self, obj):
|
||||
return obj
|
||||
|
|
@ -303,3 +304,50 @@ async def test_update_spend_logs_multiple_batches_with_failure():
|
|||
|
||||
# Verify all logs were cleared from transactions
|
||||
assert len(prisma_client.spend_log_transactions) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_insert_spend_log_queue_cap_drops_new_payload():
|
||||
"""When queue is at capacity, _insert_spend_log_to_db should drop new payloads."""
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
|
||||
writer = DBSpendUpdateWriter()
|
||||
prisma_client = MockPrismaClient()
|
||||
prisma_client.spend_log_transactions = [
|
||||
{"id": "1", "spend": 10},
|
||||
{"id": "2", "spend": 20},
|
||||
]
|
||||
payload = {"id": "3", "request_id": "req-3", "spend": 30}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.db_spend_update_writer.MAX_SPEND_LOG_QUEUE_SIZE", 2
|
||||
):
|
||||
await writer._insert_spend_log_to_db(
|
||||
payload=payload,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert len(prisma_client.spend_log_transactions) == 2
|
||||
assert payload not in prisma_client.spend_log_transactions
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_insert_spend_log_queue_cap_allows_append_below_limit():
|
||||
"""When queue has capacity, _insert_spend_log_to_db should append payload."""
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
|
||||
writer = DBSpendUpdateWriter()
|
||||
prisma_client = MockPrismaClient()
|
||||
prisma_client.spend_log_transactions = [{"id": "1", "spend": 10}]
|
||||
payload = {"id": "2", "request_id": "req-2", "spend": 20}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.db_spend_update_writer.MAX_SPEND_LOG_QUEUE_SIZE", 2
|
||||
):
|
||||
await writer._insert_spend_log_to_db(
|
||||
payload=payload,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
assert len(prisma_client.spend_log_transactions) == 2
|
||||
assert prisma_client.spend_log_transactions[-1] == payload
|
||||
|
|
|
|||
|
|
@ -88,6 +88,36 @@ def test_max_connections_url_config_none_value():
|
|||
assert pool.max_connections == 50
|
||||
|
||||
|
||||
def test_max_connections_url_config_clamped_by_soft_cap(monkeypatch):
|
||||
"""Excessive max_connections values should be clamped by soft cap."""
|
||||
monkeypatch.setenv("REDIS_MAX_CONNECTIONS_SOFT_CAP", "100")
|
||||
with patch("litellm._redis._get_redis_client_logic") as mock_logic:
|
||||
mock_logic.return_value = {
|
||||
"url": "redis://localhost:6379/0",
|
||||
"max_connections": "2147483648",
|
||||
}
|
||||
|
||||
pool = get_redis_connection_pool()
|
||||
|
||||
assert pool.max_connections == 100
|
||||
|
||||
|
||||
def test_max_connections_non_url_config_clamped_by_soft_cap(monkeypatch):
|
||||
"""Soft cap should apply to non-URL Redis connection pools too."""
|
||||
monkeypatch.setenv("REDIS_MAX_CONNECTIONS_SOFT_CAP", "120")
|
||||
with patch("litellm._redis._get_redis_client_logic") as mock_logic:
|
||||
mock_logic.return_value = {
|
||||
"host": "localhost",
|
||||
"port": 6379,
|
||||
"db": 0,
|
||||
"max_connections": "2000",
|
||||
}
|
||||
|
||||
pool = get_redis_connection_pool()
|
||||
|
||||
assert pool.max_connections == 120
|
||||
|
||||
|
||||
def _make_redis_cache():
|
||||
"""Create a RedisCache with all external I/O mocked out."""
|
||||
mock_sync_client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -329,6 +329,56 @@ class TestProxyInitializationHelpers:
|
|||
call_args = mock_uvicorn_run.call_args
|
||||
assert call_args[1]["timeout_keep_alive"] == 30
|
||||
|
||||
@patch("builtins.print")
|
||||
def test_keepalive_timeout_flag_passed_to_gunicorn(self, mock_print):
|
||||
"""Test that keepalive_timeout is passed through to gunicorn path."""
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
mock_app = MagicMock()
|
||||
mock_proxy_config = MagicMock()
|
||||
mock_key_mgmt = MagicMock()
|
||||
mock_save_worker_config = MagicMock()
|
||||
|
||||
clean_env = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in ("DATABASE_URL", "DIRECT_URL")
|
||||
}
|
||||
with patch.dict(
|
||||
os.environ, clean_env, clear=True,
|
||||
), patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"proxy_server": MagicMock(
|
||||
app=mock_app,
|
||||
ProxyConfig=mock_proxy_config,
|
||||
KeyManagementSettings=mock_key_mgmt,
|
||||
save_worker_config=mock_save_worker_config,
|
||||
)
|
||||
},
|
||||
), patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
|
||||
) as mock_get_args, patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._run_gunicorn_server"
|
||||
) as mock_run_gunicorn:
|
||||
mock_get_args.return_value = {
|
||||
"app": "litellm.proxy.proxy_server:app",
|
||||
"host": "localhost",
|
||||
"port": 8000,
|
||||
}
|
||||
|
||||
result = runner.invoke(
|
||||
run_server, ["--local", "--run_gunicorn", "--keepalive_timeout", "37"]
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
|
||||
mock_run_gunicorn.assert_called_once()
|
||||
assert mock_run_gunicorn.call_args.kwargs["keepalive_timeout"] == 37
|
||||
|
||||
@patch("uvicorn.run")
|
||||
@patch("builtins.print")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue