mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): default max_idle_connection_lifetime to prevent stale PostgreSQL connections
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
eb0e3f8c18
commit
8741d33a9b
5 changed files with 148 additions and 0 deletions
|
|
@ -2415,6 +2415,16 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
database_connection_timeout: float | None = Field(
|
||||
60, description="default timeout for a connection to the database"
|
||||
)
|
||||
database_connection_idle_lifetime: float | None = Field(
|
||||
60,
|
||||
description=(
|
||||
"Prisma `max_idle_connection_lifetime` URL param (seconds). Connections "
|
||||
"idle longer than this are closed by the pool before a managed database "
|
||||
"(RDS, Cloud SQL, Azure) silently drops them, preventing intermittent "
|
||||
"`Error { kind: Closed }` failures. Set to null to fall back to "
|
||||
"Prisma's built-in default (300s)."
|
||||
),
|
||||
)
|
||||
database_connect_timeout: float | None = Field(
|
||||
None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -82,6 +82,7 @@ CONNECTION_PARAM_KEYS: Final[frozenset[str]] = frozenset(
|
|||
"pool_timeout",
|
||||
"connect_timeout",
|
||||
"socket_timeout",
|
||||
"max_idle_connection_lifetime",
|
||||
"pgbouncer",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ telemetry: Final = None
|
|||
class LiteLLMDatabaseConnectionPool(Enum):
|
||||
database_connection_pool_limit = 10
|
||||
database_connection_pool_timeout = 60
|
||||
database_connection_idle_lifetime = 60
|
||||
|
||||
|
||||
def _build_db_connection_url_params(
|
||||
|
|
@ -63,6 +64,7 @@ def _build_db_connection_url_params(
|
|||
pool_timeout: float | None,
|
||||
connect_timeout: float | None = None,
|
||||
socket_timeout: float | None = None,
|
||||
idle_connection_lifetime: float | None = None,
|
||||
disable_prepared_statements: bool = False,
|
||||
extra_params: dict | None = None,
|
||||
) -> dict:
|
||||
|
|
@ -86,6 +88,8 @@ def _build_db_connection_url_params(
|
|||
params["connect_timeout"] = connect_timeout
|
||||
if socket_timeout is not None:
|
||||
params["socket_timeout"] = socket_timeout
|
||||
if idle_connection_lifetime is not None:
|
||||
params["max_idle_connection_lifetime"] = idle_connection_lifetime
|
||||
if disable_prepared_statements:
|
||||
params["pgbouncer"] = "true"
|
||||
if extra_params:
|
||||
|
|
@ -1081,6 +1085,9 @@ def run_server(
|
|||
db_connection_timeout: int | float | None = 60
|
||||
db_connect_timeout: int | float | None = None
|
||||
db_socket_timeout: int | float | None = None
|
||||
db_connection_idle_lifetime: int | float | None = (
|
||||
LiteLLMDatabaseConnectionPool.database_connection_idle_lifetime.value
|
||||
)
|
||||
db_disable_prepared_statements: bool = False
|
||||
db_extra_connection_params: dict | None = None
|
||||
db_statement_timeout: float | None = None
|
||||
|
|
@ -1183,6 +1190,10 @@ def run_server(
|
|||
db_connection_timeout = LiteLLMDatabaseConnectionPool.database_connection_pool_timeout.value
|
||||
db_connect_timeout = general_settings.get("database_connect_timeout")
|
||||
db_socket_timeout = general_settings.get("database_socket_timeout")
|
||||
db_connection_idle_lifetime = general_settings.get(
|
||||
"database_connection_idle_lifetime",
|
||||
LiteLLMDatabaseConnectionPool.database_connection_idle_lifetime.value,
|
||||
)
|
||||
_disable_prepared_statements: Final = general_settings.get("database_disable_prepared_statements", False)
|
||||
if isinstance(_disable_prepared_statements, str):
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
|
@ -1250,6 +1261,7 @@ def run_server(
|
|||
pool_timeout=db_connection_timeout,
|
||||
connect_timeout=db_connect_timeout,
|
||||
socket_timeout=db_socket_timeout,
|
||||
idle_connection_lifetime=db_connection_idle_lifetime,
|
||||
disable_prepared_statements=db_disable_prepared_statements,
|
||||
extra_params=db_extra_connection_params,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -739,3 +739,12 @@ def test_unsupported_db_scheme_message_names_var_and_scheme():
|
|||
assert "DIRECT_URL" in msg
|
||||
assert "sqlite" in msg
|
||||
assert "postgresql://" in msg
|
||||
|
||||
|
||||
def test_reader_shareable_params_includes_idle_lifetime():
|
||||
from litellm.proxy.db.db_url_settings import reader_shareable_params
|
||||
|
||||
shared = reader_shareable_params(
|
||||
{"max_idle_connection_lifetime": 60, "schema": "other", "connection_limit": 10}
|
||||
)
|
||||
assert shared == {"max_idle_connection_lifetime": 60, "connection_limit": 10}
|
||||
|
|
|
|||
|
|
@ -879,6 +879,122 @@ class TestProxyInitializationHelpers:
|
|||
assert appended_params["pgbouncer"] == "true"
|
||||
assert appended_params["statement_cache_size"] == 0
|
||||
|
||||
def test_build_db_connection_url_params_includes_idle_lifetime(self):
|
||||
from litellm.proxy.proxy_cli import _build_db_connection_url_params
|
||||
|
||||
params = _build_db_connection_url_params(
|
||||
connection_limit=10,
|
||||
pool_timeout=60,
|
||||
idle_connection_lifetime=60,
|
||||
)
|
||||
assert params["max_idle_connection_lifetime"] == 60
|
||||
|
||||
def test_build_db_connection_url_params_omits_none_idle_lifetime(self):
|
||||
from litellm.proxy.proxy_cli import _build_db_connection_url_params
|
||||
|
||||
params = _build_db_connection_url_params(
|
||||
connection_limit=10,
|
||||
pool_timeout=60,
|
||||
idle_connection_lifetime=None,
|
||||
)
|
||||
assert "max_idle_connection_lifetime" not in params
|
||||
|
||||
def test_build_db_connection_url_params_extra_overrides_idle_lifetime(self):
|
||||
from litellm.proxy.proxy_cli import _build_db_connection_url_params
|
||||
|
||||
params = _build_db_connection_url_params(
|
||||
connection_limit=10,
|
||||
pool_timeout=60,
|
||||
idle_connection_lifetime=60,
|
||||
extra_params={"max_idle_connection_lifetime": 300},
|
||||
)
|
||||
assert params["max_idle_connection_lifetime"] == 300
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"general_settings, expected_idle_lifetime",
|
||||
[
|
||||
({}, 60),
|
||||
({"database_connection_idle_lifetime": 30}, 30),
|
||||
],
|
||||
)
|
||||
@patch("subprocess.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
@patch(
|
||||
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
|
||||
)
|
||||
def test_db_connection_idle_lifetime_forwarded_to_url(
|
||||
self,
|
||||
mock_should_update,
|
||||
mock_setup_db,
|
||||
mock_atexit_register,
|
||||
mock_subprocess_run,
|
||||
general_settings,
|
||||
expected_idle_lifetime,
|
||||
):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
runner = CliRunner()
|
||||
mock_subprocess_run.return_value = MagicMock(returncode=0)
|
||||
|
||||
mock_proxy_module = MagicMock(
|
||||
app=MagicMock(),
|
||||
ProxyConfig=MagicMock(),
|
||||
KeyManagementSettings=MagicMock(),
|
||||
save_worker_config=MagicMock(),
|
||||
)
|
||||
mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock(
|
||||
return_value={
|
||||
"general_settings": {
|
||||
"database_url": "postgresql://test:test@localhost:5432/test",
|
||||
**general_settings,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
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": mock_proxy_module,
|
||||
"litellm.proxy.proxy_server": mock_proxy_module,
|
||||
},
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
|
||||
) as mock_get_args,
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.append_query_params",
|
||||
side_effect=lambda url, params: str(url),
|
||||
) as mock_append_query_params,
|
||||
):
|
||||
mock_get_args.return_value = {
|
||||
"app": "litellm.proxy.proxy_server:app",
|
||||
"host": "localhost",
|
||||
"port": 8000,
|
||||
}
|
||||
|
||||
result = runner.invoke(
|
||||
run_server,
|
||||
["--local", "--config", "test-config.yaml", "--skip_server_startup"],
|
||||
)
|
||||
|
||||
assert (
|
||||
result.exit_code == 0
|
||||
), f"exit_code={result.exit_code}, output={result.output}"
|
||||
mock_append_query_params.assert_called()
|
||||
appended_params = mock_append_query_params.call_args.args[1]
|
||||
assert appended_params["max_idle_connection_lifetime"] == expected_idle_lifetime
|
||||
|
||||
def test_build_db_connection_url_params_disable_prepared_statements(self):
|
||||
from litellm.proxy.proxy_cli import _build_db_connection_url_params
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue