diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py index 37d965e40a0..d393aa1b977 100644 --- a/litellm/proxy/db/db_url_settings.py +++ b/litellm/proxy/db/db_url_settings.py @@ -34,6 +34,7 @@ password when their ``*_READ_REPLICA`` counterpart is unset. import os import urllib.parse +from functools import partial from typing import Annotated, Final, cast from pydantic import AliasChoices, BeforeValidator, Field @@ -50,7 +51,10 @@ from litellm.proxy.db.token_auth import ( token_auth_flag_enabled, ) -TokenAuthFlag = Annotated[bool, BeforeValidator(token_auth_flag_enabled)] +IamTokenAuthFlag = Annotated[bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=IAM_TOKEN_DB_AUTH_ENV_VAR))] +AzureTokenAuthFlag = Annotated[ + bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR)) +] # schema.prisma pins `provider = "postgresql"`, so these are the only schemes # Prisma can actually connect with. @@ -98,8 +102,8 @@ class DatabaseURLSettings(BaseSettings): model_config = SettingsConfigDict(case_sensitive=False, extra="ignore") - iam_token_db_auth: TokenAuthFlag = Field(default=False, validation_alias=IAM_TOKEN_DB_AUTH_ENV_VAR) - azure_postgresql_auth: TokenAuthFlag = Field(default=False, validation_alias=AZURE_POSTGRESQL_AUTH_ENV_VAR) + iam_token_db_auth: IamTokenAuthFlag = Field(default=False, validation_alias=IAM_TOKEN_DB_AUTH_ENV_VAR) + azure_postgresql_auth: AzureTokenAuthFlag = Field(default=False, validation_alias=AZURE_POSTGRESQL_AUTH_ENV_VAR) # Writer database_url: str | None = Field(default=None, validation_alias="DATABASE_URL") diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index 757f7576047..fc761fc1831 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -155,6 +155,11 @@ class PrismaWrapper: # Fallback refresh interval if token parsing fails (10 minutes) FALLBACK_REFRESH_INTERVAL_SECONDS = 600 + # Floor on the proactive loop's sleep, so a token whose expiry does not advance + # (azure-identity hands back its cached token when a renewal attempt fails) costs + # one retry every 30 seconds instead of spinning the loop with no sleep at all. + TOKEN_REFRESH_MIN_SLEEP_SECONDS = 30 + ENGINE_RETIREMENT_DRAIN_TIMEOUT_SECONDS = 90 def __init__( @@ -376,8 +381,9 @@ class PrismaWrapper: For a 15-minute (900s) token with 180s buffer, this returns ~720s (12 min). Returns: - Number of seconds to sleep before the next refresh. - Returns 0 if token should be refreshed immediately. + Number of seconds to sleep before the next refresh, never less than + TOKEN_REFRESH_MIN_SLEEP_SECONDS so a token whose expiry never advances + cannot spin the loop. Returns FALLBACK_REFRESH_INTERVAL_SECONDS if parsing fails. """ db_url: Final = os.getenv(self._db_url_env_var) @@ -399,8 +405,10 @@ class PrismaWrapper: now: Final = datetime.utcnow() seconds_until_refresh: Final = (refresh_at - now).total_seconds() - # If already past refresh time, return 0 (refresh immediately) - return max(0, seconds_until_refresh) + # Past refresh time means refresh as soon as the floor allows, not instantly: + # a provider that keeps handing back the same token would otherwise leave the + # loop re-minting and recreating the query engine with no sleep between passes. + return max(self.TOKEN_REFRESH_MIN_SLEEP_SECONDS, seconds_until_refresh) def is_token_expired(self, token_url: str | None) -> bool: """Check if the token in the given URL is expired.""" diff --git a/litellm/proxy/db/token_auth.py b/litellm/proxy/db/token_auth.py index 608a07a60d5..e1f84d1c04c 100644 --- a/litellm/proxy/db/token_auth.py +++ b/litellm/proxy/db/token_auth.py @@ -36,25 +36,52 @@ CONFLICTING_TOKEN_AUTH_MESSAGE: Final = ( DEFAULT_POSTGRES_PORT: Final = "5432" TRUTHY_TOKEN_AUTH_VALUES: Final[frozenset[str]] = frozenset({"1", "on", "t", "true", "y", "yes"}) +FALSY_TOKEN_AUTH_VALUES: Final[frozenset[str]] = frozenset({"", "0", "f", "false", "n", "no", "off"}) -def token_auth_flag_enabled(value: str | bool | None) -> bool: - """Whether a token-auth toggle is on. +def token_auth_flag_enabled(value: str | bool | None, *, env_var: str) -> bool: + """Whether a token-auth toggle is on, rejecting anything it cannot read. The single parser for both toggles. Every entry point (the settings model, the CLI, and the refresh loop's own env lookup) routes through this, so a value like ``"1"`` cannot enable minting in one place and leave the refresh loop convinced token auth is off, which would strand a pod on a token it never renews. + + A value that is neither recognizably on nor recognizably off raises: silently + reading a typo as off would downgrade an operator from token auth to password + auth, and the first sign of it would be a connection refused by the server. """ if isinstance(value, bool): return value - return value is not None and value.strip().lower() in TRUTHY_TOKEN_AUTH_VALUES + if value is None: + return False + normalized: Final = value.strip().lower() + if normalized in TRUTHY_TOKEN_AUTH_VALUES: + return True + if normalized in FALSY_TOKEN_AUTH_VALUES: + return False + raise ValueError( + f"{env_var}={value!r} is not a recognized boolean. Set it to one of " + f"{', '.join(sorted(TRUTHY_TOKEN_AUTH_VALUES))} to turn token auth on, or to one of " + f"{', '.join(sorted(v for v in FALSY_TOKEN_AUTH_VALUES if v))} to turn it off." + ) def _quote(value: str) -> str: return urllib.parse.quote(value, safe="") +def _normalize_quote(value: str) -> str: + """Percent-encode a URL component that may already be percent-encoded. + + ``DATABASE_USER`` used to be interpolated raw, so pre-encoding was the only way to + put an ``@`` in it. Encoding such a value again would double-escape it, so decode + first: the round trip is idempotent and leaves an already-encoded value byte for + byte as it was, while a raw UPN like ``svc@corp`` still comes out encoded. + """ + return urllib.parse.quote(urllib.parse.unquote(value), safe="") + + @dataclass(frozen=True, slots=True) class IAMEndpoint: """Static parts of a token-authenticated Postgres connection. @@ -74,14 +101,18 @@ class IAMEndpoint: def build_url(self, token: str) -> str: """Assemble the connection URL, inserting ``token`` verbatim as the password. - User, database name, and schema are percent-encoded because an Entra principal - is a UPN containing ``@``. The token is not: both providers hand it back already - in wire form, and re-encoding it would double-escape the password. + User, database name, and schema are normalized rather than encoded outright, + because an Entra principal is a UPN containing ``@`` while an operator on the + older RDS path may already have encoded that ``@`` themselves. The token is + left alone: both providers hand it back already in wire form, and re-encoding + it would double-escape the password. """ - base: Final = f"postgresql://{_quote(self.user)}:{token}@{self.host}:{self.port}/{_quote(self.name)}" + base: Final = ( + f"postgresql://{_normalize_quote(self.user)}:{token}@{self.host}:{self.port}/{_normalize_quote(self.name)}" + ) if not self.schema: return base - return f"{base}?schema={_quote(self.schema)}" + return f"{base}?schema={_normalize_quote(self.schema)}" def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint: @@ -234,6 +265,10 @@ def build_database_token_auth(*, iam_token_db_auth: bool, azure_postgresql_auth: def resolve_database_token_auth() -> DatabaseTokenAuth | None: """Resolve the token strategy from the environment, raising when both toggles are set.""" return build_database_token_auth( - iam_token_db_auth=token_auth_flag_enabled(os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR)), - azure_postgresql_auth=token_auth_flag_enabled(os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR)), + iam_token_db_auth=token_auth_flag_enabled( + os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR + ), + azure_postgresql_auth=token_auth_flag_enabled( + os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR), env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR + ), ) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 8ad20cf1c5f..0e3e43accef 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1095,9 +1095,11 @@ def run_server( token_auth_flag_enabled, ) - wants_rds_iam: Final = iam_token_db_auth or token_auth_flag_enabled(os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR)) + wants_rds_iam: Final = iam_token_db_auth or token_auth_flag_enabled( + os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR + ) wants_azure_entra: Final = azure_postgresql_auth or token_auth_flag_enabled( - os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR) + os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR), env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR ) if wants_rds_iam: os.environ[IAM_TOKEN_DB_AUTH_ENV_VAR] = "True" diff --git a/tests/test_litellm/proxy/db/test_db_url_settings.py b/tests/test_litellm/proxy/db/test_db_url_settings.py index 35ee8349f90..e83e8310626 100644 --- a/tests/test_litellm/proxy/db/test_db_url_settings.py +++ b/tests/test_litellm/proxy/db/test_db_url_settings.py @@ -15,6 +15,7 @@ import os from unittest.mock import patch import pytest +from pydantic import ValidationError from litellm.proxy.db.db_url_settings import ( DatabaseURLSettings, @@ -114,6 +115,35 @@ def test_assembles_writer_url_when_iam_enabled(monkeypatch): assert "DATABASE_URL_READ_REPLICA" not in os.environ +def test_a_pre_encoded_iam_user_survives_url_assembly(monkeypatch): + """This URL used to be interpolated raw, so pre-encoding ``DATABASE_USER`` was the + only way to run IAM auth as a user whose name contains an ``@``. Encoding it again + yields ``svc%2540corp``, which Postgres rejects with + ``User `svc%40corp` was denied access``.""" + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true") + monkeypatch.setenv("DATABASE_HOST", "writer.example.com") + monkeypatch.setenv("DATABASE_USER", "svc%40corp") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with _stub_iam_token("WRITER_TOKEN"): + assert _apply() is True + + assert os.environ["DATABASE_URL"] == "postgresql://svc%40corp:WRITER_TOKEN@writer.example.com:5432/litellm_db" + + +def test_an_unreadable_toggle_fails_the_settings_model(monkeypatch): + """Pydantic rejected `IAM_TOKEN_DB_AUTH=enabled` before token auth had its own + parser. Reading it as 'off' instead would silently drop an operator who asked for + token auth down to password auth, with no log line saying so.""" + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "enabled") + monkeypatch.setenv("DATABASE_HOST", "writer.example.com") + monkeypatch.setenv("DATABASE_USER", "litellm") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with pytest.raises(ValidationError, match="IAM_TOKEN_DB_AUTH"): + DatabaseURLSettings.from_env() + + def test_missing_writer_envs_raises(monkeypatch): monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true") # DATABASE_HOST intentionally unset. diff --git a/tests/test_litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py index a67f48f1a74..395f17e85ef 100644 --- a/tests/test_litellm/proxy/db/test_prisma_client.py +++ b/tests/test_litellm/proxy/db/test_prisma_client.py @@ -273,6 +273,22 @@ def test_azure_entra_refresh_is_scheduled_off_the_jwt_expiry(azure_env): assert expected - 5 <= seconds <= expected +def test_a_token_whose_expiry_never_advances_cannot_spin_the_refresh_loop(azure_env): + """azure-identity hands back its cached token when a renewal attempt fails inside its + own window, so a transient Entra or IMDS problem in the last 3 minutes of a token + yields a successful refresh whose `exp` has not moved. With no floor on the sleep the + loop then re-mints and recreates the query engine on every pass, with nothing in + between, for as long as Entra stays sick.""" + wrapper = _azure_wrapper(_entra_jwt(60)) + wrapper.get_rds_iam_token() + first = wrapper._calculate_seconds_until_refresh() + + wrapper.get_rds_iam_token() + second = wrapper._calculate_seconds_until_refresh() + + assert first == second == PrismaWrapper.TOKEN_REFRESH_MIN_SLEEP_SECONDS + + def test_azure_entra_token_expiry_is_detected(azure_env): wrapper = _azure_wrapper(_entra_jwt(3600)) fresh_url = wrapper.get_rds_iam_token() diff --git a/tests/test_litellm/proxy/db/test_token_auth.py b/tests/test_litellm/proxy/db/test_token_auth.py index 493daf5ac23..56bdcd6f3e3 100644 --- a/tests/test_litellm/proxy/db/test_token_auth.py +++ b/tests/test_litellm/proxy/db/test_token_auth.py @@ -160,6 +160,25 @@ def test_build_url_encodes_a_upn_user_and_the_schema(): ) +@pytest.mark.parametrize( + ("field", "value"), + [ + ("user", "svc%40corp"), + ("name", "litellm%20db"), + ("schema", "app%2Fschema"), + ], +) +def test_build_url_leaves_an_already_encoded_component_alone(field, value): + """RDS IAM auth interpolated these raw, so pre-encoding was the only way to get an + ``@`` into ``DATABASE_USER``. Encoding again turns ``svc%40corp`` into + ``svc%2540corp``, which Postgres rejects with ``User `svc%40corp` was denied + access``, so an operator who did that on RDS breaks on upgrade.""" + url = _endpoint(**{field: value}).build_url("TOKEN") + + assert value in url + assert "%25" not in url + + def test_build_url_inserts_the_token_verbatim(): """Both providers hand the token back already in wire form, so re-encoding it here would double-escape the password.""" @@ -263,8 +282,8 @@ def test_every_truthy_spelling_enables_token_auth(monkeypatch, value): assert isinstance(resolve_database_token_auth(), AzureEntraTokenAuth) -@pytest.mark.parametrize("value", ["", " ", "false", "False", "0", "no", "off", "maybe"]) -def test_falsy_and_unrecognized_spellings_leave_token_auth_off(monkeypatch, value): +@pytest.mark.parametrize("value", ["", " ", "false", "False", "0", "no", "off", "F", "N"]) +def test_falsy_spellings_leave_token_auth_off(monkeypatch, value): """An empty string is how a Kubernetes manifest spells 'off'.""" monkeypatch.setenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, value) monkeypatch.setenv(IAM_TOKEN_DB_AUTH_ENV_VAR, value) @@ -272,6 +291,24 @@ def test_falsy_and_unrecognized_spellings_leave_token_auth_off(monkeypatch, valu assert resolve_database_token_auth() is None +@pytest.mark.parametrize("env_var", [IAM_TOKEN_DB_AUTH_ENV_VAR, AZURE_POSTGRESQL_AUTH_ENV_VAR]) +@pytest.mark.parametrize("value", ["enabled", "maybe", "TRUEE", "2"]) +def test_an_unreadable_toggle_is_a_startup_error(monkeypatch, env_var, value): + """Reading a typo as 'off' would quietly downgrade an operator who asked for token + auth to password auth, and the first sign of it is the server refusing the + connection. Pydantic rejected these before token auth had its own parser.""" + monkeypatch.setenv(env_var, value) + monkeypatch.delenv( + AZURE_POSTGRESQL_AUTH_ENV_VAR if env_var == IAM_TOKEN_DB_AUTH_ENV_VAR else IAM_TOKEN_DB_AUTH_ENV_VAR, + raising=False, + ) + + with pytest.raises(ValueError, match=env_var) as raised: + resolve_database_token_auth() + + assert value in str(raised.value) + + def test_the_entra_provider_is_built_once_per_process(): """Each build is another Azure credential with its own transport and token cache that nothing closes, and the writer, the reader, and the refresh loop each ask."""