diff --git a/litellm/_redis.py b/litellm/_redis.py index c61582abd1a..d74437c40ef 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -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: diff --git a/litellm/constants.py b/litellm/constants.py index 3d2cebf2224..d638f10bdf8 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 0c25424ceaa..4659857cfb4 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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 diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index f5163114983..b2a5b5c63ee 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -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: diff --git a/tests/proxy_unit_tests/test_update_spend.py b/tests/proxy_unit_tests/test_update_spend.py index 3734dfc5d51..1910e3e61e2 100644 --- a/tests/proxy_unit_tests/test_update_spend.py +++ b/tests/proxy_unit_tests/test_update_spend.py @@ -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 diff --git a/tests/test_litellm/caching/test_redis_connection_pool.py b/tests/test_litellm/caching/test_redis_connection_pool.py index f6e429ceff9..60bcd40eb58 100644 --- a/tests/test_litellm/caching/test_redis_connection_pool.py +++ b/tests/test_litellm/caching/test_redis_connection_pool.py @@ -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() diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index c6b2015984e..039d0f21f5c 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -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")