Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/token_auth.py: 41%
131 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Token-based authentication for the proxy's Postgres connection.
3Two managed Postgres offerings hand the client a short-lived credential that is used as
4the Postgres password: AWS RDS with IAM auth, and Azure Database for PostgreSQL Flexible
5Server with Microsoft Entra ID. Both need the same machinery (mint at startup, read the
6expiry back off the token, mint again before it lapses) and differ only in how the token
7is produced and how its expiry is encoded, so the difference lives in a tagged union that
8is resolved once from the environment and injected into whatever needs a token.
9"""
11import base64
12import functools
13import os
14import urllib.parse
15from collections.abc import Callable
16from dataclasses import dataclass
17from datetime import datetime, timedelta, timezone
18from typing import Final, TypeAlias
20from pydantic import BaseModel
21from typing_extensions import assert_never
23from litellm._logging import verbose_proxy_logger
25IAM_TOKEN_DB_AUTH_ENV_VAR: Final = "IAM_TOKEN_DB_AUTH"
26AZURE_POSTGRESQL_AUTH_ENV_VAR: Final = "AZURE_POSTGRESQL_AUTH"
27AZURE_POSTGRESQL_SCOPE: Final = "https://ossrdbms-aad.database.windows.net/.default"
29CONFLICTING_TOKEN_AUTH_MESSAGE: Final = (
30 f"{IAM_TOKEN_DB_AUTH_ENV_VAR} and {AZURE_POSTGRESQL_AUTH_ENV_VAR} are both enabled, but the "
31 "database password can only come from one token source. Keep "
32 f"{IAM_TOKEN_DB_AUTH_ENV_VAR} for AWS RDS IAM auth, or {AZURE_POSTGRESQL_AUTH_ENV_VAR} for "
33 "Azure Database for PostgreSQL with Microsoft Entra ID, and unset the other one."
34)
36DEFAULT_POSTGRES_PORT: Final = "5432"
38TRUTHY_TOKEN_AUTH_VALUES: Final[frozenset[str]] = frozenset({"1", "on", "t", "true", "y", "yes"})
39FALSY_TOKEN_AUTH_VALUES: Final[frozenset[str]] = frozenset({"", "0", "f", "false", "n", "no", "off"})
42def token_auth_flag_enabled(value: str | bool | None, *, env_var: str) -> bool:
43 """Whether a token-auth toggle is on, rejecting anything it cannot read.
45 The single parser for both toggles. Every entry point (the settings model, the
46 CLI, and the refresh loop's own env lookup) routes through this, so a value like
47 ``"1"`` cannot enable minting in one place and leave the refresh loop convinced
48 token auth is off, which would strand a pod on a token it never renews.
50 A value that is neither recognizably on nor recognizably off raises: silently
51 reading a typo as off would downgrade an operator from token auth to password
52 auth, and the first sign of it would be a connection refused by the server.
53 """
54 if isinstance(value, bool): 54 ↛ 55line 54 didn't jump to line 55 because the condition on line 54 was never true
55 return value
56 if value is None: 56 ↛ 58line 56 didn't jump to line 58 because the condition on line 56 was always true
57 return False
58 normalized: Final = value.strip().lower()
59 if normalized in TRUTHY_TOKEN_AUTH_VALUES:
60 return True
61 if normalized in FALSY_TOKEN_AUTH_VALUES:
62 return False
63 raise ValueError(
64 f"{env_var}={value!r} is not a recognized boolean. Set it to one of "
65 f"{', '.join(sorted(TRUTHY_TOKEN_AUTH_VALUES))} to turn it on, or to one of "
66 f"{', '.join(sorted(v for v in FALSY_TOKEN_AUTH_VALUES if v))} to turn it off."
67 )
70def _quote(value: str) -> str:
71 return urllib.parse.quote(value, safe="")
74def _normalize_quote(value: str) -> str:
75 """Percent-encode a URL component that may already be percent-encoded.
77 ``DATABASE_USER`` used to be interpolated raw, so pre-encoding was the only way to
78 put an ``@`` in it. Encoding such a value again would double-escape it, so decode
79 first: the round trip is idempotent and leaves an already-encoded value byte for
80 byte as it was, while a raw UPN like ``svc@corp`` still comes out encoded.
81 """
82 return urllib.parse.quote(urllib.parse.unquote(value), safe="")
85@dataclass(frozen=True, slots=True)
86class IAMEndpoint:
87 """Static parts of a token-authenticated Postgres connection.
89 The token rotates every few minutes to an hour depending on the provider;
90 everything else (host, port, user, database name, schema) stays fixed. Capturing
91 the static fields once means a refresh only regenerates the token and reassembles
92 the URL.
93 """
95 host: str
96 port: str
97 user: str
98 name: str
99 schema: str | None = None
101 def build_url(self, token: str) -> str:
102 """Assemble the connection URL, inserting ``token`` verbatim as the password.
104 User, database name, and schema are normalized rather than encoded outright,
105 because an Entra principal is a UPN containing ``@`` while an operator on the
106 older RDS path may already have encoded that ``@`` themselves. The token is
107 left alone: both providers hand it back already in wire form, and re-encoding
108 it would double-escape the password.
109 """
110 base: Final = (
111 f"postgresql://{_normalize_quote(self.user)}:{token}@{self.host}:{self.port}/{_normalize_quote(self.name)}"
112 )
113 if not self.schema:
114 return base
115 return f"{base}?schema={_normalize_quote(self.schema)}"
118def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint:
119 """Parse an :class:`IAMEndpoint` back out of a Postgres URL.
121 Used so a reader URL can drive its own token refresh without requiring callers to
122 set parallel ``DATABASE_HOST_READ_REPLICA`` / etc. env vars.
123 """
124 parsed: Final = urllib.parse.urlparse(url)
125 if not parsed.hostname or not parsed.username:
126 raise ValueError("Cannot parse IAM endpoint from URL: missing host or username")
127 name: Final = urllib.parse.unquote((parsed.path or "/").lstrip("/"))
128 if not name:
129 raise ValueError("Cannot parse IAM endpoint from URL: missing database name")
130 port: Final = str(parsed.port) if parsed.port else DEFAULT_POSTGRES_PORT
131 schema_values: Final = urllib.parse.parse_qs(parsed.query).get("schema") if parsed.query else None
132 return IAMEndpoint(
133 host=parsed.hostname,
134 port=port,
135 user=urllib.parse.unquote(parsed.username),
136 name=name,
137 schema=schema_values[0] if schema_values else None,
138 )
141@dataclass(frozen=True, slots=True)
142class RdsIamTokenAuth:
143 """AWS RDS IAM auth: a SigV4-presigned token minted from the ambient AWS credentials."""
145 @property
146 def label(self) -> str:
147 return "RDS IAM token"
149 @property
150 def env_var(self) -> str:
151 return IAM_TOKEN_DB_AUTH_ENV_VAR
154@dataclass(frozen=True, slots=True)
155class AzureEntraTokenAuth:
156 """Azure Database for PostgreSQL auth: a Microsoft Entra ID access token as the password.
158 The provider is injected rather than resolved here so callers (and tests) decide which
159 Azure credential mints the token.
160 """
162 token_provider: Callable[[], str]
164 @property
165 def label(self) -> str:
166 return "Azure Entra token"
168 @property
169 def env_var(self) -> str:
170 return AZURE_POSTGRESQL_AUTH_ENV_VAR
173DatabaseTokenAuth: TypeAlias = RdsIamTokenAuth | AzureEntraTokenAuth
176def mint_database_token(auth: DatabaseTokenAuth, endpoint: IAMEndpoint) -> str:
177 """Mint a fresh database password for ``endpoint``, already percent-encoded."""
178 match auth:
179 case RdsIamTokenAuth():
180 from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token
182 return generate_iam_auth_token(db_host=endpoint.host, db_port=endpoint.port, db_user=endpoint.user)
183 case AzureEntraTokenAuth():
184 return _quote(auth.token_provider())
185 case _:
186 assert_never(auth)
189def parse_database_token_expiration(auth: DatabaseTokenAuth, token: str) -> datetime | None:
190 """Return when ``token`` expires as a naive UTC datetime, or None when unreadable.
192 Callers fall back to a fixed refresh interval on None, so an unparseable token
193 degrades to periodic refresh instead of failing.
194 """
195 match auth:
196 case RdsIamTokenAuth():
197 return _parse_rds_token_expiration(token)
198 case AzureEntraTokenAuth():
199 return _parse_entra_token_expiration(token)
200 case _:
201 assert_never(auth)
204def _parse_rds_token_expiration(token: str) -> datetime | None:
205 if "?" not in token:
206 return None
207 try:
208 params: Final = urllib.parse.parse_qs(token.split("?", 1)[1])
209 expires_values: Final = params.get("X-Amz-Expires")
210 date_values: Final = params.get("X-Amz-Date")
211 if not expires_values or not date_values:
212 return None
213 created: Final = datetime.strptime(date_values[0], "%Y%m%dT%H%M%SZ")
214 return created + timedelta(seconds=int(expires_values[0]))
215 except (ValueError, OverflowError, OSError) as exc:
216 verbose_proxy_logger.debug("Failed to parse RDS IAM token expiration: %s", exc)
217 return None
220class _EntraAccessTokenClaims(BaseModel):
221 exp: int
224def _parse_entra_token_expiration(token: str) -> datetime | None:
225 segments: Final = token.split(".")
226 if len(segments) != 3:
227 return None
228 payload: Final = segments[1]
229 try:
230 claims: Final = _EntraAccessTokenClaims.model_validate_json(
231 base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4))
232 )
233 except ValueError as exc:
234 verbose_proxy_logger.debug("Failed to parse Azure Entra token expiration: %s", exc)
235 return None
236 return datetime.fromtimestamp(claims.exp, tz=timezone.utc).replace(tzinfo=None)
239@functools.cache
240def build_azure_entra_token_provider() -> Callable[[], str]:
241 """The process-wide Entra token provider for the Azure Postgres OSS RDBMS scope.
243 Cached because the writer URL, the reader URL, and the refresh loop each ask for a
244 strategy, and every uncached call would build another Azure credential with its own
245 HTTP transport and its own token cache that nothing ever closes.
246 """
247 from litellm.secret_managers.get_azure_ad_token_provider import (
248 get_azure_ad_token_provider,
249 )
251 return get_azure_ad_token_provider(azure_scope=AZURE_POSTGRESQL_SCOPE)
254def build_database_token_auth(*, iam_token_db_auth: bool, azure_postgresql_auth: bool) -> DatabaseTokenAuth | None:
255 """Pick the token strategy the two toggles ask for, or None when neither is on."""
256 if iam_token_db_auth and azure_postgresql_auth: 256 ↛ 257line 256 didn't jump to line 257 because the condition on line 256 was never true
257 raise RuntimeError(CONFLICTING_TOKEN_AUTH_MESSAGE)
258 if azure_postgresql_auth: 258 ↛ 259line 258 didn't jump to line 259 because the condition on line 258 was never true
259 return AzureEntraTokenAuth(token_provider=build_azure_entra_token_provider())
260 if iam_token_db_auth: 260 ↛ 261line 260 didn't jump to line 261 because the condition on line 260 was never true
261 return RdsIamTokenAuth()
262 return None
265def resolve_database_token_auth() -> DatabaseTokenAuth | None:
266 """Resolve the token strategy from the environment, raising when both toggles are set."""
267 return build_database_token_auth(
268 iam_token_db_auth=token_auth_flag_enabled(
269 os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR
270 ),
271 azure_postgresql_auth=token_auth_flag_enabled(
272 os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR), env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR
273 ),
274 )