mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
50bf99a909
commit
4eb7bf32c0
7 changed files with 153 additions and 21 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue