Add queue and redis caps; pass gunicorn keepalive

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-02-27 05:11:50 +00:00
parent fd58c8c060
commit 78c2a39f1c
7 changed files with 215 additions and 16 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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