diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 3acc19d397d..2062ca93fb3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -78,6 +78,23 @@ class _InvalidIndex: table_size: str MAX_MIGRATE_DEPLOY_ATTEMPTS = 4 +LIBPQ_URL_PARAMS: Final = frozenset( + { + "sslmode", + "sslcert", + "sslkey", + "sslrootcert", + "sslpassword", + "application_name", + "connect_timeout", + "client_encoding", + "options", + "service", + "gssencmode", + "krbsrvname", + "target_session_attrs", + } +) @dataclass(frozen=True) @@ -689,30 +706,43 @@ class ProxyExtrasDBManager: @staticmethod def _strip_prisma_query_params(url: str) -> str: - """Remove Prisma-specific query params (connection_limit, pool_timeout, - schema, etc.) from DATABASE_URL so psycopg can parse it.""" + """Rewrite a Prisma-dialect URL for libpq: drop the Prisma-only params + (connection_limit, pool_timeout, schema, pgbouncer, sslaccept, ...) and + translate Prisma's TLS params back, since libpq reads ``sslcert`` as a + client certificate where Prisma reads it as the CA.""" from urllib.parse import parse_qsl, quote, urlencode, urlparse, urlunparse - parsed = urlparse(url) + parsed: Final = urlparse(url) if not parsed.query: return url - libpq_params = { - "sslmode", - "sslcert", - "sslkey", - "sslrootcert", - "sslpassword", - "application_name", - "connect_timeout", - "client_encoding", - "options", - "service", - "gssencmode", - "krbsrvname", - "target_session_attrs", - } - kept = [(k, v) for k, v in parse_qsl(parsed.query) if k in libpq_params] - return urlunparse(parsed._replace(query=urlencode(kept, quote_via=quote))) + pairs: Final = tuple(parse_qsl(parsed.query)) + kept: Final = tuple((k, v) for k, v in pairs if k in LIBPQ_URL_PARAMS) + sslaccept: Final = next((v for k, v in pairs if k == "sslaccept"), None) + libpq_pairs: Final = ProxyExtrasDBManager._libpq_tls_params(kept, sslaccept) + return urlunparse(parsed._replace(query=urlencode(libpq_pairs, quote_via=quote))) + + @staticmethod + def _libpq_tls_params( + pairs: "tuple[tuple[str, str], ...]", sslaccept: "str | None" + ) -> "tuple[tuple[str, str], ...]": + """Undo ``translate_libpq_ssl_params``. Prisma's ``sslcert`` is the CA and + ``sslaccept=strict`` checks chain and hostname, which libpq only does in + ``sslmode=verify-full``, so strict becomes ``sslrootcert`` plus + ``verify-full`` whatever ``sslmode`` said (``disable`` stays off). Prisma + defaults an absent ``sslaccept`` to ``accept_invalid_certs`` and anything + else to strict. Without strict it checks nothing, so the CA is dropped and + ``sslmode`` is kept as is: libpq only verifies when a root cert is present. + A URL that also carries ``sslkey`` is libpq's own client-certificate form + and is kept.""" + keys: Final = frozenset(k for k, _ in pairs) + if "sslcert" not in keys or "sslkey" in keys: + return pairs + sslmode: Final = next((v for k, v in pairs if k == "sslmode"), None) + rest: Final = tuple((k, v) for k, v in pairs if k not in ("sslcert", "sslmode")) + if sslaccept in (None, "accept_invalid_certs") or sslmode == "disable": + return rest if sslmode is None else rest + (("sslmode", sslmode),) + root_cert: Final = tuple(("sslrootcert", v) for k, v in pairs if k == "sslcert" and "sslrootcert" not in keys) + return rest + root_cert + (("sslmode", "verify-full"),) @staticmethod def _warn_if_db_ahead_of_head(migrations_dir: str) -> None: diff --git a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py index 5540cf54193..29c9ec56d91 100644 --- a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py +++ b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py @@ -1029,6 +1029,84 @@ class TestJWTKeyMappingCascade: +class TestStripPrismaQueryParams: + """The psycopg URL the job connects with is derived from the Prisma-dialect + DATABASE_URL, whose TLS params mean something else to libpq.""" + + @staticmethod + def _query(url: str) -> dict[str, str]: + from urllib.parse import parse_qsl, urlparse + + return dict(parse_qsl(urlparse(url).query)) + + def test_prisma_ca_sslcert_becomes_sslrootcert_with_verify_full(self): + url = "postgresql://u:p@writer:5432/db?schema=public&sslmode=require&sslcert=/tmp/pinned.pem&sslaccept=strict" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/tmp/pinned.pem"} + assert cleaned.startswith("postgresql://u:p@writer:5432/db?") + + @pytest.mark.parametrize("sslmode", ["prefer", "require"]) + @pytest.mark.parametrize("sslaccept", ["strict", "unknown-mode-prisma-treats-as-strict"]) + def test_strict_verifies_chain_and_hostname_whatever_sslmode_prisma_was_given(self, sslmode, sslaccept): + url = f"postgresql://writer/db?sslmode={sslmode}&sslcert=/certs/ca.pem&sslaccept={sslaccept}" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/certs/ca.pem"} + + def test_strict_with_tls_disabled_stays_off(self): + url = "postgresql://writer/db?sslmode=disable&sslcert=/certs/ca.pem&sslaccept=strict" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "disable"} + + @pytest.mark.parametrize("sslaccept", ["&sslaccept=accept_invalid_certs", ""]) + def test_without_strict_the_ca_is_dropped_so_libpq_checks_nothing_like_prisma(self, sslaccept): + url = f"postgresql://writer/db?sslmode=require&sslcert=/certs/ca.pem{sslaccept}" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "require"} + + def test_a_ca_alone_without_strict_or_sslmode_leaves_libpq_its_defaults(self): + cleaned = ProxyExtrasDBManager._strip_prisma_query_params("postgresql://writer/db?sslcert=/certs/ca.pem") + + assert cleaned == "postgresql://writer/db" + + def test_a_libpq_client_certificate_pair_is_left_alone(self): + url = "postgresql://writer/db?sslmode=verify-full&sslrootcert=/ca.pem&sslcert=/client.crt&sslkey=/client.key" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == { + "sslmode": "verify-full", + "sslrootcert": "/ca.pem", + "sslcert": "/client.crt", + "sslkey": "/client.key", + } + + def test_an_explicit_sslrootcert_wins_over_the_prisma_sslcert(self): + url = "postgresql://writer/db?sslmode=require&sslrootcert=/ca.pem&sslcert=/pinned.pem&sslaccept=strict" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert self._query(cleaned) == {"sslmode": "verify-full", "sslrootcert": "/ca.pem"} + + def test_prisma_only_params_are_dropped_and_plain_urls_pass_through(self): + url = "postgresql://u:p@pooler:6543/db?schema=tenant&pgbouncer=true&connection_limit=5&connect_timeout=3" + + cleaned = ProxyExtrasDBManager._strip_prisma_query_params(url) + + assert cleaned == "postgresql://u:p@pooler:6543/db?connect_timeout=3" + assert ( + ProxyExtrasDBManager._strip_prisma_query_params("postgresql://u:p@writer/db") + == "postgresql://u:p@writer/db" + ) + + class TestBuildRequestLogIndexes: """The migration job hands the index build the direct database URL and the schema the migrations target, waits for it, and reports its result."""