mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: fail fast for non-Postgres database URLs (#30883)
* fix(proxy): fail fast on non-PostgreSQL DATABASE_URL instead of hanging on startup LiteLLM's Prisma datasource is pinned to provider = 'postgresql', so a sqlite:// or mysql:// DATABASE_URL can never connect. Today that surfaces as an opaque startup stall where the port never binds, and a separate 'DB not connected' 500 on /key/generate when no DATABASE_URL is set at all leaves operators guessing what to configure. Validate the DATABASE_URL / DIRECT_URL scheme in run_server before any Prisma call and exit with an actionable message naming the unsupported scheme. Also reword CommonProxyErrors.db_not_connected_error to tell the operator to set DATABASE_URL to a postgresql:// connection string. Add regression tests covering postgres acceptance and sqlite/mysql/mssql rejection. * fix: resolve CI failures and proxy DB URL typing issue * fix(proxy): fail fast on non-PostgreSQL DATABASE_URLs with clear startup errors instead of hanging * Validate DIRECT_URL alongside DATABASE_URL startup guards
This commit is contained in:
parent
19450bdb9d
commit
488f7874c6
5 changed files with 238 additions and 20 deletions
|
|
@ -3353,7 +3353,9 @@ class ProxyException(Exception):
|
|||
|
||||
class CommonProxyErrors(str, enum.Enum):
|
||||
db_not_connected_error = (
|
||||
"DB not connected. See https://docs.litellm.ai/docs/proxy/virtual_keys"
|
||||
"DB not connected. This endpoint needs a database; set DATABASE_URL to a "
|
||||
"PostgreSQL connection string (postgresql://...) to enable it. "
|
||||
"See https://docs.litellm.ai/docs/proxy/virtual_keys"
|
||||
)
|
||||
no_llm_router = "No models configured on proxy"
|
||||
not_allowed_access = "Admin-only endpoint. Not allowed to access this."
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ password when their ``*_READ_REPLICA`` counterpart is unset.
|
|||
|
||||
import os
|
||||
import urllib.parse
|
||||
from typing import Optional, cast
|
||||
from typing import Final, cast
|
||||
|
||||
from pydantic import AliasChoices, Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
|
@ -44,6 +44,41 @@ from litellm.proxy.auth import rds_iam_token
|
|||
_IAM_ENV_KEY = "IAM_TOKEN_DB_AUTH"
|
||||
_DEFAULT_PG_PORT = "5432"
|
||||
|
||||
# schema.prisma pins `provider = "postgresql"`, so these are the only schemes
|
||||
# Prisma can actually connect with.
|
||||
SUPPORTED_DB_SCHEMES: Final[frozenset[str]] = frozenset({"postgresql", "postgres"})
|
||||
_MISSING_SCHEME = "<missing scheme>"
|
||||
|
||||
|
||||
def unsupported_db_scheme(database_url: str) -> str | None:
|
||||
"""Return the connection URL scheme when it is not PostgreSQL, else None.
|
||||
|
||||
A `sqlite://` / `mysql://` URL can never connect against the
|
||||
postgresql-only datasource, but the resulting Prisma failure is opaque and
|
||||
version-dependent (a confusing migration error, or a startup that never
|
||||
binds). Callers use this to reject the URL up front with an actionable
|
||||
error instead.
|
||||
|
||||
A schemeless value (e.g. a malformed DSN like ``user:pass@host/db``) yields
|
||||
the ``_MISSING_SCHEME`` placeholder rather than the raw URL, so callers that
|
||||
log the return value never echo embedded credentials.
|
||||
"""
|
||||
scheme = urllib.parse.urlsplit(database_url).scheme.lower()
|
||||
if scheme in SUPPORTED_DB_SCHEMES:
|
||||
return None
|
||||
return scheme or _MISSING_SCHEME
|
||||
|
||||
|
||||
def unsupported_db_scheme_message(env_var: str, scheme: str) -> str:
|
||||
"""Operator-facing message naming the offending env var and scheme."""
|
||||
return (
|
||||
f"{env_var} uses unsupported scheme '{scheme}'. LiteLLM's database "
|
||||
"features (virtual keys, store_model_in_db, spend tracking) require "
|
||||
"PostgreSQL; use a 'postgresql://' connection string. SQLite and other "
|
||||
"engines are not supported. "
|
||||
"See https://docs.litellm.ai/docs/proxy/virtual_keys"
|
||||
)
|
||||
|
||||
|
||||
class DatabaseURLSettings(BaseSettings):
|
||||
"""Discrete ``DATABASE_*`` env vars, loaded once at process start.
|
||||
|
|
@ -58,46 +93,47 @@ class DatabaseURLSettings(BaseSettings):
|
|||
iam_token_db_auth: bool = Field(default=False, validation_alias=_IAM_ENV_KEY)
|
||||
|
||||
# Writer
|
||||
database_url: Optional[str] = Field(default=None, validation_alias="DATABASE_URL")
|
||||
database_host: Optional[str] = Field(default=None, validation_alias="DATABASE_HOST")
|
||||
database_url: str | None = Field(default=None, validation_alias="DATABASE_URL")
|
||||
direct_url: str | None = Field(default=None, validation_alias="DIRECT_URL")
|
||||
database_host: str | None = Field(default=None, validation_alias="DATABASE_HOST")
|
||||
database_port: str = Field(
|
||||
default=_DEFAULT_PG_PORT, validation_alias="DATABASE_PORT"
|
||||
)
|
||||
database_user: Optional[str] = Field(
|
||||
database_user: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("DATABASE_USER", "DATABASE_USERNAME"),
|
||||
)
|
||||
database_name: Optional[str] = Field(default=None, validation_alias="DATABASE_NAME")
|
||||
database_schema: Optional[str] = Field(
|
||||
database_name: str | None = Field(default=None, validation_alias="DATABASE_NAME")
|
||||
database_schema: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_SCHEMA"
|
||||
)
|
||||
database_password: Optional[str] = Field(
|
||||
database_password: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_PASSWORD"
|
||||
)
|
||||
|
||||
# Read replica
|
||||
database_url_read_replica: Optional[str] = Field(
|
||||
database_url_read_replica: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_URL_READ_REPLICA"
|
||||
)
|
||||
database_host_read_replica: Optional[str] = Field(
|
||||
database_host_read_replica: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_HOST_READ_REPLICA"
|
||||
)
|
||||
database_port_read_replica: Optional[str] = Field(
|
||||
database_port_read_replica: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_PORT_READ_REPLICA"
|
||||
)
|
||||
database_user_read_replica: Optional[str] = Field(
|
||||
database_user_read_replica: str | None = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices(
|
||||
"DATABASE_USER_READ_REPLICA", "DATABASE_USERNAME_READ_REPLICA"
|
||||
),
|
||||
)
|
||||
database_name_read_replica: Optional[str] = Field(
|
||||
database_name_read_replica: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_NAME_READ_REPLICA"
|
||||
)
|
||||
database_schema_read_replica: Optional[str] = Field(
|
||||
database_schema_read_replica: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_SCHEMA_READ_REPLICA"
|
||||
)
|
||||
database_password_read_replica: Optional[str] = Field(
|
||||
database_password_read_replica: str | None = Field(
|
||||
default=None, validation_alias="DATABASE_PASSWORD_READ_REPLICA"
|
||||
)
|
||||
|
||||
|
|
@ -106,7 +142,7 @@ class DatabaseURLSettings(BaseSettings):
|
|||
"""Load the settings from ``os.environ`` (read at call time)."""
|
||||
return cls()
|
||||
|
||||
def build_writer_url(self) -> Optional[str]:
|
||||
def build_writer_url(self) -> str | None:
|
||||
"""Return the writer URL to set, or ``None`` to leave it as-is.
|
||||
|
||||
Raises ``RuntimeError`` (naming the offending vars) when IAM auth is
|
||||
|
|
@ -156,7 +192,7 @@ class DatabaseURLSettings(BaseSettings):
|
|||
)
|
||||
return None
|
||||
|
||||
def build_reader_url(self) -> Optional[str]:
|
||||
def build_reader_url(self) -> str | None:
|
||||
"""Return the read-replica URL to set, or ``None`` to leave it as-is.
|
||||
|
||||
Opt-in via ``DATABASE_HOST_READ_REPLICA``; never clobbers a
|
||||
|
|
@ -217,11 +253,11 @@ class DatabaseURLSettings(BaseSettings):
|
|||
def _password_url(
|
||||
*,
|
||||
user: str,
|
||||
password: Optional[str],
|
||||
password: str | None,
|
||||
host: str,
|
||||
port: str,
|
||||
name: str,
|
||||
schema: Optional[str],
|
||||
schema: str | None,
|
||||
) -> str:
|
||||
"""Percent-encode credentials into a ``postgresql://`` URL.
|
||||
|
||||
|
|
@ -239,6 +275,26 @@ class DatabaseURLSettings(BaseSettings):
|
|||
url += f"?schema={schema}"
|
||||
return url
|
||||
|
||||
def _raise_for_unsupported_scheme(self) -> None:
|
||||
"""Reject an operator-pinned non-PostgreSQL writer / direct / reader URL.
|
||||
|
||||
The componentized entrypoints (gateway / backend / migrations) call
|
||||
``apply_to_env`` and then hand the URL straight to Prisma, bypassing
|
||||
the CLI's own guard. A pinned URL flows through untouched, so validate
|
||||
the same three vars the CLI guard checks (DATABASE_URL, DIRECT_URL, and
|
||||
the read replica) rather than letting Prisma stall on an unusable scheme.
|
||||
"""
|
||||
for env_var, url in (
|
||||
("DATABASE_URL", self.database_url),
|
||||
("DIRECT_URL", self.direct_url),
|
||||
("DATABASE_URL_READ_REPLICA", self.database_url_read_replica),
|
||||
):
|
||||
if not url:
|
||||
continue
|
||||
bad_scheme = unsupported_db_scheme(url)
|
||||
if bad_scheme is not None:
|
||||
raise RuntimeError(unsupported_db_scheme_message(env_var, bad_scheme))
|
||||
|
||||
def apply_to_env(self) -> bool:
|
||||
"""Write the assembled URL(s) into ``os.environ``.
|
||||
|
||||
|
|
@ -246,6 +302,7 @@ class DatabaseURLSettings(BaseSettings):
|
|||
password auth that assembled a fresh URL). False means there was
|
||||
nothing to do — an operator-pinned URL, or no discrete fields.
|
||||
"""
|
||||
self._raise_for_unsupported_scheme()
|
||||
wrote_writer = False
|
||||
writer_url = self.build_writer_url()
|
||||
if writer_url is not None:
|
||||
|
|
|
|||
|
|
@ -1195,6 +1195,25 @@ def run_server(
|
|||
os.getenv("DATABASE_URL", None) is not None
|
||||
or os.getenv("DIRECT_URL", None) is not None
|
||||
):
|
||||
from litellm.proxy.db.db_url_settings import (
|
||||
unsupported_db_scheme,
|
||||
unsupported_db_scheme_message,
|
||||
)
|
||||
|
||||
for _db_env in ("DATABASE_URL", "DIRECT_URL"):
|
||||
_candidate_url = os.getenv(_db_env)
|
||||
if _candidate_url is None:
|
||||
continue
|
||||
_bad_scheme = unsupported_db_scheme(_candidate_url)
|
||||
if _bad_scheme is not None:
|
||||
print(
|
||||
f"\033[1;31mLiteLLM Proxy: "
|
||||
f"{unsupported_db_scheme_message(_db_env, _bad_scheme)}"
|
||||
"\033[0m",
|
||||
file=sys.stderr,
|
||||
flush=True,
|
||||
)
|
||||
sys.exit(1)
|
||||
try:
|
||||
from litellm.secret_managers.main import get_secret
|
||||
|
||||
|
|
|
|||
|
|
@ -16,7 +16,11 @@ from unittest.mock import patch
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.db.db_url_settings import DatabaseURLSettings
|
||||
from litellm.proxy.db.db_url_settings import (
|
||||
DatabaseURLSettings,
|
||||
unsupported_db_scheme,
|
||||
unsupported_db_scheme_message,
|
||||
)
|
||||
|
||||
|
||||
def _apply() -> bool:
|
||||
|
|
@ -27,6 +31,7 @@ def _apply() -> bool:
|
|||
_MANAGED_DB_ENV_VARS = (
|
||||
"IAM_TOKEN_DB_AUTH",
|
||||
"DATABASE_URL",
|
||||
"DIRECT_URL",
|
||||
"DATABASE_URL_READ_REPLICA",
|
||||
"DATABASE_HOST",
|
||||
"DATABASE_PORT",
|
||||
|
|
@ -287,3 +292,87 @@ def test_password_reader_uses_own_credentials(monkeypatch):
|
|||
os.environ["DATABASE_URL_READ_REPLICA"]
|
||||
== "postgresql://litellm_ro:ro_pw@reader.example.com:5432/litellm_db"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url",
|
||||
[
|
||||
"postgresql://u:p@host:5432/db",
|
||||
"postgres://u:p@host:5432/db",
|
||||
"POSTGRESQL://u:p@host:5432/db",
|
||||
"postgresql://host/db?schema=public",
|
||||
],
|
||||
)
|
||||
def test_unsupported_db_scheme_accepts_postgres(url):
|
||||
assert unsupported_db_scheme(url) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url,scheme",
|
||||
[
|
||||
("sqlite:///data/litellm.db", "sqlite"),
|
||||
("sqlite:///./local.db", "sqlite"),
|
||||
("mysql://u:p@host:3306/db", "mysql"),
|
||||
("mssql://host/db", "mssql"),
|
||||
],
|
||||
)
|
||||
def test_unsupported_db_scheme_rejects_non_postgres(url, scheme):
|
||||
assert unsupported_db_scheme(url) == scheme
|
||||
|
||||
|
||||
def test_unsupported_db_scheme_does_not_echo_schemeless_credentials():
|
||||
"""A malformed schemeless DSN must not leak its embedded credentials
|
||||
through the return value (which callers log)."""
|
||||
leaky = "litellm:s3cr3t_password@db.internal:5432/litellm"
|
||||
|
||||
result = unsupported_db_scheme(leaky)
|
||||
|
||||
assert result is not None
|
||||
assert "s3cr3t_password" not in result
|
||||
assert "db.internal" not in result
|
||||
|
||||
|
||||
def test_apply_to_env_rejects_pinned_sqlite_writer(monkeypatch):
|
||||
"""Componentized entrypoints pin DATABASE_URL and call apply_to_env; a
|
||||
sqlite writer must raise here rather than reach Prisma."""
|
||||
monkeypatch.setenv("DATABASE_URL", "sqlite:///data/litellm.db")
|
||||
|
||||
with pytest.raises(RuntimeError, match="sqlite"):
|
||||
_apply()
|
||||
|
||||
# The bad URL must not have been propagated as a usable connection string.
|
||||
assert os.environ["DATABASE_URL"] == "sqlite:///data/litellm.db"
|
||||
|
||||
|
||||
def test_apply_to_env_rejects_pinned_sqlite_direct_url(monkeypatch):
|
||||
"""DIRECT_URL reaches Prisma the same way DATABASE_URL does; a non-postgres
|
||||
direct URL must be rejected in apply_to_env, matching the CLI startup guard."""
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
|
||||
monkeypatch.setenv("DIRECT_URL", "sqlite:///data/litellm.db")
|
||||
|
||||
with pytest.raises(RuntimeError, match="DIRECT_URL.*sqlite"):
|
||||
_apply()
|
||||
|
||||
|
||||
def test_apply_to_env_rejects_pinned_non_postgres_reader(monkeypatch):
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db"
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="DATABASE_URL_READ_REPLICA.*mysql"):
|
||||
_apply()
|
||||
|
||||
|
||||
def test_apply_to_env_accepts_pinned_postgres(monkeypatch):
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@host:5432/db")
|
||||
|
||||
# Operator-pinned URL: nothing reassembled, no error.
|
||||
assert _apply() is False
|
||||
|
||||
|
||||
def test_unsupported_db_scheme_message_names_var_and_scheme():
|
||||
msg = unsupported_db_scheme_message("DIRECT_URL", "sqlite")
|
||||
assert "DIRECT_URL" in msg
|
||||
assert "sqlite" in msg
|
||||
assert "postgresql://" in msg
|
||||
|
|
|
|||
|
|
@ -1708,6 +1708,57 @@ class TestRunServerDbSetup:
|
|||
use_migrate=True, use_v2_resolver=False
|
||||
)
|
||||
|
||||
@patch("subprocess.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
@patch("litellm.proxy.db.check_migration.check_prisma_schema_diff")
|
||||
@patch("litellm.proxy.db.prisma_client.should_update_prisma_schema")
|
||||
def test_startup_exits_on_non_postgres_database_url(
|
||||
self,
|
||||
mock_should_update_schema,
|
||||
mock_check_schema_diff,
|
||||
mock_setup_database,
|
||||
mock_atexit_register,
|
||||
mock_subprocess_run,
|
||||
):
|
||||
"""A sqlite DATABASE_URL must exit immediately, before any prisma call,
|
||||
instead of stalling on a migration against the postgresql-only schema."""
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
mock_subprocess_run.return_value = MagicMock(returncode=0)
|
||||
mock_should_update_schema.return_value = True
|
||||
|
||||
mock_proxy_module = MagicMock(
|
||||
app=MagicMock(),
|
||||
ProxyConfig=MagicMock(),
|
||||
KeyManagementSettings=MagicMock(),
|
||||
save_worker_config=MagicMock(),
|
||||
)
|
||||
|
||||
clean_env = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in ("DATABASE_URL", "DIRECT_URL")
|
||||
}
|
||||
clean_env["DATABASE_URL"] = "sqlite:///data/litellm.db"
|
||||
|
||||
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,
|
||||
},
|
||||
),
|
||||
):
|
||||
with pytest.raises(SystemExit) as exc_info:
|
||||
run_server.main(
|
||||
["--local", "--skip_server_startup"], standalone_mode=False
|
||||
)
|
||||
assert exc_info.value.code == 1
|
||||
mock_setup_database.assert_not_called()
|
||||
|
||||
|
||||
# --- Module-level helpers for worker startup hook tests ---
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue