fix(proxy): keep token-auth URLs, toggles, and refresh sleeps safe

Three fixes on the Postgres token-auth path found by a live risk pass:

Pre-encoded connection components no longer double-escape. The user, database
name, and schema used to be interpolated raw, so encoding an already-encoded
DATABASE_USER like svc%40corp turned it into svc%2540corp and Postgres rejected
the login with P1010. Decoding before encoding is idempotent, so a pre-encoded
value comes out byte for byte as it went in while a raw UPN still gets encoded.

An unreadable IAM_TOKEN_DB_AUTH or AZURE_POSTGRESQL_AUTH now fails startup
naming the variable and the value. Reading a typo like "enabled" as off would
silently downgrade an operator from token auth to password auth, and the first
sign of it would be the server refusing the connection.

The proactive refresh loop floors its sleep at 30 seconds. azure-identity hands
back its cached token when a renewal fails inside its own window, so a token
whose expiry never advances used to compute a zero sleep and spin the loop,
re-minting and recreating the Prisma query engine every pass.

Co-authored-by: David Balatoni <balcsida@gmail.com>
This commit is contained in:
mateo-berri 2026-08-20 12:59:58 -07:00
parent 50bf99a909
commit 4eb7bf32c0
7 changed files with 153 additions and 21 deletions

View file

@ -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")

View file

@ -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."""

View file

@ -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
),
)

View file

@ -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"

View file

@ -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.

View file

@ -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()

View file

@ -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."""