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

1"""Token-based authentication for the proxy's Postgres connection. 

2 

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

10 

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 

19 

20from pydantic import BaseModel 

21from typing_extensions import assert_never 

22 

23from litellm._logging import verbose_proxy_logger 

24 

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" 

28 

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) 

35 

36DEFAULT_POSTGRES_PORT: Final = "5432" 

37 

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

40 

41 

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. 

44 

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. 

49 

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 ) 

68 

69 

70def _quote(value: str) -> str: 

71 return urllib.parse.quote(value, safe="") 

72 

73 

74def _normalize_quote(value: str) -> str: 

75 """Percent-encode a URL component that may already be percent-encoded. 

76 

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

83 

84 

85@dataclass(frozen=True, slots=True) 

86class IAMEndpoint: 

87 """Static parts of a token-authenticated Postgres connection. 

88 

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

94 

95 host: str 

96 port: str 

97 user: str 

98 name: str 

99 schema: str | None = None 

100 

101 def build_url(self, token: str) -> str: 

102 """Assemble the connection URL, inserting ``token`` verbatim as the password. 

103 

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

116 

117 

118def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint: 

119 """Parse an :class:`IAMEndpoint` back out of a Postgres URL. 

120 

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 ) 

139 

140 

141@dataclass(frozen=True, slots=True) 

142class RdsIamTokenAuth: 

143 """AWS RDS IAM auth: a SigV4-presigned token minted from the ambient AWS credentials.""" 

144 

145 @property 

146 def label(self) -> str: 

147 return "RDS IAM token" 

148 

149 @property 

150 def env_var(self) -> str: 

151 return IAM_TOKEN_DB_AUTH_ENV_VAR 

152 

153 

154@dataclass(frozen=True, slots=True) 

155class AzureEntraTokenAuth: 

156 """Azure Database for PostgreSQL auth: a Microsoft Entra ID access token as the password. 

157 

158 The provider is injected rather than resolved here so callers (and tests) decide which 

159 Azure credential mints the token. 

160 """ 

161 

162 token_provider: Callable[[], str] 

163 

164 @property 

165 def label(self) -> str: 

166 return "Azure Entra token" 

167 

168 @property 

169 def env_var(self) -> str: 

170 return AZURE_POSTGRESQL_AUTH_ENV_VAR 

171 

172 

173DatabaseTokenAuth: TypeAlias = RdsIamTokenAuth | AzureEntraTokenAuth 

174 

175 

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 

181 

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) 

187 

188 

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. 

191 

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) 

202 

203 

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 

218 

219 

220class _EntraAccessTokenClaims(BaseModel): 

221 exp: int 

222 

223 

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) 

237 

238 

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. 

242 

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 ) 

250 

251 return get_azure_ad_token_provider(azure_scope=AZURE_POSTGRESQL_SCOPE) 

252 

253 

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 

263 

264 

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 )