Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/prisma_client.py: 39%

442 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2This file contains the PrismaWrapper class, which wraps the Prisma client and keeps the 

3database token (AWS RDS IAM or Microsoft Entra ID) fresh. 

4""" 

5 

6import asyncio 

7import os 

8import random 

9import signal 

10import subprocess 

11import time 

12import urllib 

13import urllib.parse 

14from collections.abc import Callable 

15from datetime import datetime, timedelta 

16from typing import TYPE_CHECKING, Any, Final, Protocol 

17 

18from litellm._logging import verbose_proxy_logger 

19from litellm.proxy.db.db_url_settings import add_missing_query_params, token_refresh_params_from_url 

20from litellm.proxy.db.token_auth import ( 

21 DEFAULT_POSTGRES_PORT, 

22 DatabaseTokenAuth, 

23 IAMEndpoint, 

24 RdsIamTokenAuth, 

25 mint_database_token, 

26 parse_database_token_expiration, 

27 parse_iam_endpoint_from_url, 

28) 

29from litellm.secret_managers.main import str_to_bool 

30 

31if TYPE_CHECKING: 31 ↛ 32line 31 didn't jump to line 32 because the condition on line 31 was never true

32 from prisma import Prisma 

33 

34__all__ = ( 

35 "IAMEndpoint", 

36 "PrismaManager", 

37 "PrismaWrapper", 

38 "parse_iam_endpoint_from_url", 

39) 

40 

41 

42class _PrismaProcess(Protocol): 

43 pid: object 

44 

45 

46class _PrismaEngine(Protocol): 

47 @property 

48 def process(self) -> _PrismaProcess: ... 48 ↛ exitline 48 didn't return from function 'process' because

49 

50 async def query(self, content: str, *, tx_id: str | None) -> object: ... 50 ↛ exitline 50 didn't return from function 'query' because

51 

52 async def start_transaction(self, *, content: str) -> str: ... 52 ↛ exitline 52 didn't return from function 'start_transaction' because

53 

54 async def commit_transaction(self, tx_id: str) -> None: ... 54 ↛ exitline 54 didn't return from function 'commit_transaction' because

55 

56 async def rollback_transaction(self, tx_id: str) -> None: ... 56 ↛ exitline 56 didn't return from function 'rollback_transaction' because

57 

58 

59class _PrismaClient(Protocol): 

60 _Prisma__engine: _PrismaEngine 

61 

62 @property 

63 def _engine(self) -> _PrismaEngine: ... 63 ↛ exitline 63 didn't return from function '_engine' because

64 

65 

66class _PrismaDrainTracker: 

67 def __init__(self) -> None: 

68 self._active_operations = 0 

69 self._transactions: frozenset[str] = frozenset() 

70 self._drained = asyncio.Event() 

71 self._drained.set() 

72 

73 def begin_operation(self) -> None: 

74 if self._active_operations == 0: 

75 self._drained.clear() 

76 self._active_operations += 1 

77 

78 def end_operation(self) -> None: 

79 self._active_operations -= 1 

80 if self._active_operations == 0: 

81 self._drained.set() 

82 

83 def transaction_started(self, transaction_id: str) -> None: 

84 self._transactions = self._transactions.union((transaction_id,)) 

85 

86 def transaction_finished(self, transaction_id: str) -> None: 

87 if transaction_id not in self._transactions: 87 ↛ 88line 87 didn't jump to line 88 because the condition on line 87 was never true

88 return 

89 self._transactions = self._transactions.difference((transaction_id,)) 

90 self.end_operation() 

91 

92 async def wait_until_drained(self) -> None: 

93 await self._drained.wait() 

94 

95 

96class _TrackedPrismaEngine: 

97 def __init__(self, engine: _PrismaEngine, tracker: _PrismaDrainTracker) -> None: 

98 self._engine = engine 

99 self.tracker = tracker 

100 

101 @property 

102 def process(self) -> _PrismaProcess: 

103 return self._engine.process 

104 

105 def __getattr__(self, name: str) -> object: 

106 return getattr(self._engine, name) 

107 

108 async def query(self, content: str, *, tx_id: str | None) -> object: 

109 self.tracker.begin_operation() 

110 try: 

111 return await self._engine.query(content, tx_id=tx_id) 

112 finally: 

113 self.tracker.end_operation() 

114 

115 async def start_transaction(self, *, content: str) -> str: 

116 self.tracker.begin_operation() 

117 try: 

118 transaction_id: Final = await self._engine.start_transaction(content=content) 

119 except (Exception, asyncio.CancelledError): 

120 self.tracker.end_operation() 

121 raise 

122 self.tracker.transaction_started(transaction_id) 

123 return transaction_id 

124 

125 async def commit_transaction(self, tx_id: str) -> None: 

126 self.tracker.begin_operation() 

127 try: 

128 await self._engine.commit_transaction(tx_id) 

129 finally: 

130 self.tracker.end_operation() 

131 self.tracker.transaction_finished(tx_id) 

132 

133 async def rollback_transaction(self, tx_id: str) -> None: 

134 self.tracker.begin_operation() 

135 try: 

136 await self._engine.rollback_transaction(tx_id) 

137 finally: 

138 self.tracker.end_operation() 

139 self.tracker.transaction_finished(tx_id) 

140 

141 

142class PrismaWrapper: 

143 """ 

144 Wrapper around Prisma client that handles token-based database authentication. 

145 

146 When a token strategy is active (AWS RDS IAM or Microsoft Entra ID), this wrapper: 

147 1. Proactively refreshes the token before it expires (background task) 

148 2. Falls back to synchronous refresh if a token is found expired 

149 3. Uses proper locking to prevent race conditions during reconnection 

150 

151 RDS IAM tokens are valid for 15 minutes and Entra tokens for about an hour. This 

152 wrapper refreshes 3 minutes before whatever expiry the live token carries. 

153 """ 

154 

155 # Buffer time in seconds before token expiration to trigger refresh 

156 # Refresh 3 minutes (180 seconds) before the token expires 

157 TOKEN_REFRESH_BUFFER_SECONDS = 180 

158 

159 # Fallback refresh interval if token parsing fails (10 minutes) 

160 FALLBACK_REFRESH_INTERVAL_SECONDS = 600 

161 

162 # Floor on the proactive loop's sleep, so a token whose expiry does not advance 

163 # (azure-identity hands back its cached token when a renewal attempt fails) costs 

164 # one retry every 30 seconds instead of spinning the loop with no sleep at all. 

165 TOKEN_REFRESH_MIN_SLEEP_SECONDS = 30 

166 

167 ENGINE_RETIREMENT_DRAIN_TIMEOUT_SECONDS = 90 

168 

169 def __init__( 

170 self, 

171 original_prisma: Any, 

172 iam_token_db_auth: bool = False, 

173 *, 

174 token_auth: DatabaseTokenAuth | None = None, 

175 db_url_env_var: str = "DATABASE_URL", 

176 iam_endpoint: IAMEndpoint | None = None, 

177 recreate_uses_datasource: bool = False, 

178 log_prefix: str = "", 

179 ): 

180 # Set before `_original_prisma` so the `iam_token_db_auth` property below can 

181 # never send `__getattr__` looking for a half-built strategy on the raw client. 

182 self._token_auth = token_auth if token_auth is not None else (RdsIamTokenAuth() if iam_token_db_auth else None) 

183 self._original_prisma = original_prisma 

184 

185 # Per-connection knobs so the same wrapper can be used for the writer 

186 # (defaults: DATABASE_URL env, IAM endpoint from DATABASE_HOST/etc., 

187 # recreate via env reload) or for a reader (DATABASE_URL_READ_REPLICA 

188 # env, IAM endpoint parsed from that URL, recreate via datasource 

189 # override since Prisma only auto-reads DATABASE_URL). 

190 self._db_url_env_var = db_url_env_var 

191 self._iam_endpoint = iam_endpoint 

192 self._recreate_uses_datasource = recreate_uses_datasource 

193 # Tag every log line emitted by this wrapper instance so writer and 

194 # reader can be told apart in interleaved output (e.g. "[writer] RDS 

195 # IAM token refresh scheduled in 720 seconds"). Empty string (default) 

196 # keeps backward-compatible logs for the single-DB case. 

197 self._log_prefix = f"{log_prefix} " if log_prefix else "" 

198 

199 # Background token refresh task management 

200 self._token_refresh_task: asyncio.Task | None = None 

201 self._reconnection_lock = asyncio.Lock() 

202 self._last_refresh_time: datetime | None = None 

203 self._active_drain_tracker = self._instrument_prisma_client(original_prisma) 

204 self._retirement_tasks: frozenset[asyncio.Task[None]] = frozenset() 

205 

206 # Coordination for planned engine restarts (issue #29176). Every 

207 # `recreate_prisma_client` SIGTERMs the running query-engine on 

208 # purpose. The engine-death watcher (in `PrismaClient`) must be able 

209 # to tell that planned kill apart from a real crash, otherwise it 

210 # triggers its own reconnect and kills the freshly-spawned engine. 

211 # - `_expected_engine_deaths`: PIDs we intentionally killed; the 

212 # watcher consumes these instead of reconnecting. 

213 # - `_engine_generation`: monotonic counter bumped on every 

214 # successful recreate, used by callers as an optimistic-lock token 

215 # so racing/cascading recreates collapse into a single restart. 

216 # - `on_engine_replaced`: optional callback fired after a recreate so 

217 # the owner (PrismaClient) can re-arm its watcher on the new PID. 

218 self._expected_engine_deaths: set[int] = set() 

219 self._engine_generation: int = 0 

220 self.on_engine_replaced: Callable[[], None] | None = None 

221 

222 @property 

223 def token_auth(self) -> DatabaseTokenAuth | None: 

224 """The active database token strategy, or None for password auth.""" 

225 return self._token_auth 

226 

227 @property 

228 def token_label(self) -> str: 

229 """Human name of the active token kind, for log lines.""" 

230 return self._token_auth.label if self._token_auth is not None else "database token" 

231 

232 @property 

233 def iam_token_db_auth(self) -> bool: 

234 """Whether any token strategy is active. 

235 

236 Read-only: the kind of token is chosen once, by injection, so there is no way 

237 to flip this back on and silently get AWS RDS on an Azure deployment. 

238 """ 

239 return self._token_auth is not None 

240 

241 @staticmethod 

242 def _read_engine(prisma_client: _PrismaClient) -> _PrismaEngine: 

243 return prisma_client._engine 

244 

245 @staticmethod 

246 def _write_engine(prisma_client: _PrismaClient, engine: _PrismaEngine) -> None: 

247 prisma_client._Prisma__engine = engine 

248 

249 def _instrument_prisma_client(self, prisma_client: "Prisma | _PrismaClient") -> _PrismaDrainTracker | None: 

250 from prisma.errors import ClientNotConnectedError 

251 

252 try: 

253 engine: Final = self._read_engine(prisma_client) 

254 except (AttributeError, ClientNotConnectedError): 

255 return None 

256 if isinstance(engine, _TrackedPrismaEngine): 256 ↛ 257line 256 didn't jump to line 257 because the condition on line 256 was never true

257 return engine.tracker 

258 tracker: Final = _PrismaDrainTracker() 

259 self._write_engine(prisma_client, _TrackedPrismaEngine(engine, tracker)) 

260 return tracker 

261 

262 def _get_engine_pid(self, prisma_client: "Prisma | _PrismaClient | None" = None) -> int: 

263 """Get the PID of the current Prisma engine subprocess, or 0 if unavailable. 

264 

265 Must never raise: it runs inside the reconnect path, where the client 

266 may be in any broken state. Prisma's ``_engine`` is a property that 

267 raises ``ClientNotConnectedError`` on a disconnected client; if that 

268 escaped here, ``recreate_prisma_client`` would fail before it could 

269 build a replacement client and the reconnect loop could never recover. 

270 """ 

271 from prisma.errors import ClientNotConnectedError 

272 

273 try: 

274 target_prisma: Final = self._original_prisma if prisma_client is None else prisma_client 

275 pid: Final = self._read_engine(target_prisma).process.pid 

276 if isinstance(pid, int): 

277 return pid 

278 except (AttributeError, ClientNotConnectedError, TypeError): 

279 pass 

280 return 0 

281 

282 async def _retire_engine_when_drained(self, pid: int, tracker: _PrismaDrainTracker | None) -> None: 

283 if tracker is not None: 

284 try: 

285 await asyncio.wait_for( 

286 tracker.wait_until_drained(), 

287 timeout=self.ENGINE_RETIREMENT_DRAIN_TIMEOUT_SECONDS, 

288 ) 

289 except asyncio.TimeoutError: 

290 verbose_proxy_logger.warning( 

291 "%sReplaced prisma engine PID %s did not drain within %ss; killing it with work still in flight.", 

292 self._log_prefix, 

293 pid, 

294 self.ENGINE_RETIREMENT_DRAIN_TIMEOUT_SECONDS, 

295 ) 

296 await self._kill_engine_process(pid) 

297 

298 def _schedule_engine_retirement(self, pid: int, tracker: _PrismaDrainTracker | None) -> None: 

299 if pid <= 0: 

300 return 

301 retirement_task: Final = asyncio.create_task(self._retire_engine_when_drained(pid, tracker)) 

302 self._retirement_tasks = self._retirement_tasks.union((retirement_task,)) 

303 retirement_task.add_done_callback(self._retirement_finished) 

304 

305 def _retirement_finished(self, retirement_task: asyncio.Task[None]) -> None: 

306 self._retirement_tasks = self._retirement_tasks.difference((retirement_task,)) 

307 

308 async def connect(self, timeout: int | timedelta | None = None) -> None: 

309 if timeout is None: 309 ↛ 312line 309 didn't jump to line 312 because the condition on line 309 was always true

310 await self._original_prisma.connect() 

311 else: 

312 await self._original_prisma.connect(timeout) 

313 self._active_drain_tracker = self._instrument_prisma_client(self._original_prisma) 

314 

315 @staticmethod 

316 async def _kill_engine_process(pid: int) -> None: 

317 """Force-kill the engine subprocess to prevent DB connection pool leaks. 

318 

319 Called on every reconnect (in `recreate_prisma_client`) to retire the 

320 old query-engine subprocess without invoking prisma-client-py's 

321 synchronous `disconnect()` — which blocks the asyncio event loop on 

322 `subprocess.Popen.wait()` for 30-120+ seconds when the engine is 

323 stuck on TCP close. 

324 

325 Sends SIGTERM for graceful shutdown, waits briefly, then SIGKILL as 

326 a backstop. 

327 """ 

328 if pid <= 0: 

329 return 

330 try: 

331 os.kill(pid, signal.SIGTERM) 

332 except (ProcessLookupError, PermissionError, OSError): 

333 return # Already dead or inaccessible 

334 verbose_proxy_logger.warning( 

335 "Sent SIGTERM to prisma-query-engine PID %s during reconnect.", 

336 pid, 

337 ) 

338 # Brief wait for graceful shutdown, then force-kill 

339 await asyncio.sleep(0.5) 

340 try: 

341 os.kill(pid, getattr(signal, "SIGKILL", signal.SIGTERM)) 

342 verbose_proxy_logger.warning( 

343 "Sent SIGKILL to prisma-query-engine PID %s (did not exit after SIGTERM).", 

344 pid, 

345 ) 

346 except (ProcessLookupError, PermissionError, OSError): 

347 pass # Exited after SIGTERM — expected 

348 

349 def _extract_token_from_db_url(self, db_url: str | None) -> str | None: 

350 """ 

351 Extract the token (password) from the DATABASE_URL. 

352 

353 The token contains the AWS signature with X-Amz-Date and X-Amz-Expires parameters. 

354 

355 Important: We must parse the URL while it's still encoded to preserve structure, 

356 then decode the password portion. Otherwise the '?' in the token breaks URL parsing. 

357 """ 

358 if db_url is None: 

359 return None 

360 try: 

361 # Parse URL while still encoded to preserve structure 

362 parsed: Final = urllib.parse.urlparse(db_url) 

363 if parsed.password: 

364 # Now decode just the password/token 

365 return urllib.parse.unquote(parsed.password) 

366 return None 

367 except Exception: 

368 return None 

369 

370 def _parse_token_expiration(self, token: str | None) -> datetime | None: 

371 """ 

372 Parse the token to extract its expiration time. 

373 

374 Returns the datetime when the token expires, or None if parsing fails. 

375 """ 

376 if token is None or self._token_auth is None: 

377 return None 

378 return parse_database_token_expiration(self._token_auth, token) 

379 

380 def _calculate_seconds_until_refresh(self) -> float: 

381 """ 

382 Calculate exactly how many seconds until we need to refresh the token. 

383 

384 Uses precise timing: sleeps until (token_expiration - buffer_seconds). 

385 For a 15-minute (900s) token with 180s buffer, this returns ~720s (12 min). 

386 

387 Returns: 

388 Number of seconds to sleep before the next refresh, never less than 

389 TOKEN_REFRESH_MIN_SLEEP_SECONDS so a token whose expiry never advances 

390 cannot spin the loop. 

391 Returns FALLBACK_REFRESH_INTERVAL_SECONDS if parsing fails. 

392 """ 

393 db_url: Final = os.getenv(self._db_url_env_var) 

394 token: Final = self._extract_token_from_db_url(db_url) 

395 expiration_time: Final = self._parse_token_expiration(token) 

396 

397 if expiration_time is None: 

398 # If we can't parse the token, use fallback interval 

399 verbose_proxy_logger.debug( 

400 "Could not parse token expiration, using fallback interval of %ss", 

401 self.FALLBACK_REFRESH_INTERVAL_SECONDS, 

402 ) 

403 return self.FALLBACK_REFRESH_INTERVAL_SECONDS 

404 

405 # Calculate when we should refresh (expiration - buffer) 

406 refresh_at: Final = expiration_time - timedelta(seconds=self.TOKEN_REFRESH_BUFFER_SECONDS) 

407 

408 # How long until refresh time? 

409 now: Final = datetime.utcnow() 

410 seconds_until_refresh: Final = (refresh_at - now).total_seconds() 

411 

412 # Past refresh time means refresh as soon as the floor allows, not instantly: 

413 # a provider that keeps handing back the same token would otherwise leave the 

414 # loop re-minting and recreating the query engine with no sleep between passes. 

415 return max(self.TOKEN_REFRESH_MIN_SLEEP_SECONDS, seconds_until_refresh) 

416 

417 def is_token_expired(self, token_url: str | None) -> bool: 

418 """Check if the token in the given URL is expired.""" 

419 if token_url is None: 

420 return True 

421 

422 token: Final = self._extract_token_from_db_url(token_url) 

423 expiration_time: Final = self._parse_token_expiration(token) 

424 

425 if expiration_time is None: 

426 # If we can't parse the token, assume it's expired to trigger refresh 

427 verbose_proxy_logger.debug("Could not parse token expiration, treating as expired") 

428 return True 

429 

430 return datetime.utcnow() > expiration_time 

431 

432 def get_rds_iam_token(self) -> str | None: 

433 """Mint a fresh database token and update the configured DB URL env var. 

434 

435 When the wrapper was constructed with an explicit `iam_endpoint` 

436 (typical for a reader wrapper whose host/port/user came from a parsed 

437 URL), use that. Otherwise fall back to the DATABASE_HOST/PORT/USER/ 

438 NAME/SCHEMA env vars (writer behavior). 

439 """ 

440 auth: Final = self._token_auth 

441 if auth is None: 

442 return None 

443 

444 endpoint: Final = self._iam_endpoint if self._iam_endpoint is not None else self._endpoint_from_env() 

445 db_url: Final = add_missing_query_params( 

446 endpoint.build_url(mint_database_token(auth, endpoint)), 

447 token_refresh_params_from_url(os.environ.get(self._db_url_env_var, "")), 

448 ) 

449 os.environ[self._db_url_env_var] = db_url 

450 return db_url 

451 

452 @staticmethod 

453 def _endpoint_from_env() -> IAMEndpoint: 

454 host: Final = os.getenv("DATABASE_HOST") 

455 user: Final = os.getenv("DATABASE_USER") 

456 name: Final = os.getenv("DATABASE_NAME") 

457 if not host or not user or not name: 

458 missing: Final = tuple( 

459 env 

460 for env, value in (("DATABASE_HOST", host), ("DATABASE_USER", user), ("DATABASE_NAME", name)) 

461 if not value 

462 ) 

463 raise RuntimeError( 

464 f"Cannot mint a database token: {', '.join(missing)} unset. Set them so the " 

465 "connection URL can be reassembled around a freshly minted token." 

466 ) 

467 return IAMEndpoint( 

468 host=host, 

469 # Default to the Postgres standard port; passing None to 

470 # `generate_iam_auth_token` makes botocore embed the literal 

471 # string "None" in the presigned URL, which then fails to parse. 

472 port=os.getenv("DATABASE_PORT", DEFAULT_POSTGRES_PORT), 

473 user=user, 

474 name=name, 

475 schema=os.getenv("DATABASE_SCHEMA"), 

476 ) 

477 

478 @property 

479 def engine_generation(self) -> int: 

480 """How many query-engine replacements have completed on this wrapper. 

481 

482 Bumped under `_reconnection_lock` only after a replacement engine has 

483 connected, so a change across an await proves a *successful* planned 

484 replacement happened in between — a replacement that failed (a real 

485 outage) leaves it untouched. 

486 """ 

487 return self._engine_generation 

488 

489 async def _reconnection_settled(self) -> None: 

490 async with self._reconnection_lock: 

491 pass 

492 

493 async def wait_for_planned_engine_replacement(self, timeout_seconds: float) -> None: 

494 """Wait, bounded, for an in-flight planned engine replacement to finish. 

495 

496 Both replacement paths (`recreate_prisma_client` and 

497 `_safe_refresh_token`) hold `_reconnection_lock` across their whole 

498 kill/connect window, so re-acquiring it means the replacement has 

499 settled one way or the other. Gives up silently on timeout: a caller 

500 that stopped waiting must treat the replacement as not completed and 

501 consult `engine_generation` rather than assume success. 

502 """ 

503 if timeout_seconds <= 0 or not self._reconnection_lock.locked(): 

504 return 

505 try: 

506 await asyncio.wait_for(self._reconnection_settled(), timeout=timeout_seconds) 

507 except asyncio.TimeoutError: 

508 return 

509 

510 async def recreate_prisma_client( 

511 self, 

512 new_db_url: str, 

513 http_client: object | None = None, 

514 *, 

515 expected_generation: int | None = None, 

516 ) -> bool: 

517 """Disconnect and reconnect the Prisma client with a new database URL. 

518 

519 Kills the old engine subprocess directly (SIGTERM → SIGKILL) rather than 

520 calling `disconnect()`. prisma-client-py's `disconnect()` calls a 

521 synchronous `subprocess.Popen.wait()` that can freeze the asyncio event 

522 loop for 30-120+ seconds when the engine is stuck on TCP close, 

523 breaking `/health/liveliness` and causing Kubernetes pod restarts. 

524 

525 The writer wrapper relies on Prisma re-reading `DATABASE_URL` from env; 

526 the reader wrapper opts into `recreate_uses_datasource=True` so the 

527 new URL is passed explicitly via `datasource={"url": ...}` (Prisma 

528 does not auto-read alternate env vars like DATABASE_URL_READ_REPLICA). 

529 

530 Serializes all recreations through `self._reconnection_lock` so the 

531 IAM-refresh path and the engine-death/transport-error reconnect paths 

532 cannot recreate concurrently (issue #29176). `expected_generation`, if 

533 given, is an optimistic-lock token: when it no longer matches 

534 `self._engine_generation` once the lock is held, another path already 

535 replaced the engine, so this call is a no-op and returns ``False``. 

536 

537 Returns: 

538 bool: ``True`` if the client was actually recreated, ``False`` if 

539 the recreate was skipped because the engine generation moved on. 

540 """ 

541 async with self._reconnection_lock: 

542 return await self._recreate_prisma_client_locked( 

543 new_db_url, 

544 http_client=http_client, 

545 expected_generation=expected_generation, 

546 ) 

547 

548 async def _recreate_prisma_client_locked( 

549 self, 

550 new_db_url: str, 

551 http_client: object | None = None, 

552 *, 

553 expected_generation: int | None = None, 

554 ) -> bool: 

555 """Core recreate logic. Caller MUST hold `self._reconnection_lock`. 

556 

557 Split out so callers that already hold the lock (e.g. 

558 `_safe_refresh_token`, which double-checks token freshness under the 

559 lock) don't re-acquire it — `asyncio.Lock` is not reentrant. 

560 """ 

561 from prisma import Prisma 

562 

563 if expected_generation is not None and expected_generation != self._engine_generation: 

564 verbose_proxy_logger.info( 

565 "%sSkipping Prisma client recreate: engine already replaced (generation %s != expected %s).", 

566 self._log_prefix, 

567 self._engine_generation, 

568 expected_generation, 

569 ) 

570 return False 

571 

572 old_engine_pid: Final = self._get_engine_pid() 

573 if old_engine_pid > 0: 

574 # Record BEFORE the kill so the engine-death watcher, which may 

575 # fire the instant the process dies, recognizes this as a planned 

576 # restart and does not launch its own reconnect. 

577 # 

578 # A stale entry can linger when the watcher re-arms on the new PID 

579 # before the old PID's death callback runs (the callback then 

580 # early-returns on PID mismatch without consuming it). Such entries 

581 # are harmless but would accumulate on a long-running proxy (~one 

582 # per IAM refresh), so cap the set — those old PIDs are long dead. 

583 if len(self._expected_engine_deaths) >= 64: 

584 self._expected_engine_deaths.clear() 

585 self._expected_engine_deaths.add(old_engine_pid) 

586 await self._kill_engine_process(old_engine_pid) 

587 

588 kwargs: Final[dict[str, Any]] = {} 

589 if http_client is not None: 

590 kwargs["http"] = http_client 

591 if self._recreate_uses_datasource: 

592 kwargs["datasource"] = {"url": new_db_url} 

593 self._original_prisma = Prisma(**kwargs) 

594 

595 await self._original_prisma.connect() 

596 self._active_drain_tracker = self._instrument_prisma_client(self._original_prisma) 

597 self._engine_generation += 1 

598 

599 # Let the owner (PrismaClient) re-arm its engine-death watcher on the 

600 # newly-spawned engine PID. Scheduled, never awaited, so a slow watcher 

601 # can't stall the refresh while we hold the reconnection lock. 

602 if self.on_engine_replaced is not None: 

603 self.on_engine_replaced() 

604 

605 return True 

606 

607 async def _replace_prisma_client_for_token_refresh_locked( 

608 self, 

609 new_db_url: str, 

610 ) -> None: 

611 from prisma import Prisma 

612 from prisma.types import DatasourceOverride 

613 

614 if self._recreate_uses_datasource: 

615 datasource: Final = DatasourceOverride(url=new_db_url) 

616 replacement_prisma = Prisma(datasource=datasource) 

617 else: 

618 replacement_prisma = Prisma() 

619 

620 old_prisma: Final = self._original_prisma 

621 old_engine_pid: Final = self._get_engine_pid(old_prisma) 

622 old_drain_tracker = self._active_drain_tracker 

623 if old_drain_tracker is None: 

624 old_drain_tracker = self._instrument_prisma_client(old_prisma) 

625 try: 

626 await replacement_prisma.connect() 

627 except (Exception, asyncio.CancelledError): 

628 self._schedule_engine_retirement(self._get_engine_pid(replacement_prisma), None) 

629 raise 

630 replacement_drain_tracker: Final = self._instrument_prisma_client(replacement_prisma) 

631 

632 if old_engine_pid > 0: 

633 if len(self._expected_engine_deaths) >= 64: 

634 self._expected_engine_deaths.clear() 

635 self._expected_engine_deaths.add(old_engine_pid) 

636 

637 self._original_prisma = replacement_prisma 

638 self._active_drain_tracker = replacement_drain_tracker 

639 self._engine_generation += 1 

640 

641 if self.on_engine_replaced is not None: 

642 self.on_engine_replaced() 

643 

644 self._schedule_engine_retirement(old_engine_pid, old_drain_tracker) 

645 

646 async def start_token_refresh_task(self) -> None: 

647 """ 

648 Start the background token refresh task. 

649 

650 This task proactively refreshes the database token before it expires, 

651 preventing connection failures. Should be called after the initial 

652 Prisma client connection is established. 

653 """ 

654 if not self.iam_token_db_auth: 654 ↛ 658line 654 didn't jump to line 658 because the condition on line 654 was always true

655 verbose_proxy_logger.debug("Database token auth not enabled, skipping token refresh task") 

656 return 

657 

658 if self._token_refresh_task is not None: 

659 verbose_proxy_logger.debug("Token refresh task already running") 

660 return 

661 

662 self._token_refresh_task = asyncio.create_task(self._token_refresh_loop()) 

663 verbose_proxy_logger.info( 

664 "%sStarted %s proactive refresh background task", 

665 self._log_prefix, 

666 self.token_label, 

667 ) 

668 

669 async def stop_token_refresh_task(self) -> None: 

670 """ 

671 Stop the background token refresh task gracefully. 

672 

673 Should be called during application shutdown to clean up resources. 

674 """ 

675 if self._token_refresh_task is None: 675 ↛ 678line 675 didn't jump to line 678 because the condition on line 675 was always true

676 return 

677 

678 self._token_refresh_task.cancel() 

679 try: 

680 await self._token_refresh_task 

681 except asyncio.CancelledError: 

682 pass 

683 self._token_refresh_task = None 

684 verbose_proxy_logger.info( 

685 "%sStopped %s refresh background task", 

686 self._log_prefix, 

687 self.token_label, 

688 ) 

689 

690 async def _token_refresh_loop(self) -> None: 

691 """ 

692 Background loop that proactively refreshes database tokens before expiration. 

693 

694 Uses precise timing: calculates the exact sleep duration until the token 

695 needs to be refreshed (expiration - 3 minute buffer), then refreshes. 

696 This is more efficient than polling, requiring only 1 wake-up per token cycle. 

697 """ 

698 verbose_proxy_logger.info( 

699 "%s%s refresh loop started. Tokens will be refreshed %ss before expiration.", 

700 self._log_prefix, 

701 self.token_label, 

702 self.TOKEN_REFRESH_BUFFER_SECONDS, 

703 ) 

704 

705 while True: 

706 try: 

707 # Calculate exactly how long to sleep until next refresh 

708 sleep_seconds = self._calculate_seconds_until_refresh() 

709 

710 if sleep_seconds > 0: 

711 verbose_proxy_logger.info( 

712 f"{self._log_prefix}{self.token_label} refresh scheduled in " 

713 f"{sleep_seconds:.0f} seconds ({sleep_seconds / 60:.1f} minutes)" 

714 ) 

715 await asyncio.sleep(sleep_seconds) 

716 

717 # Refresh the token 

718 verbose_proxy_logger.info( 

719 "%sProactively refreshing %s...", 

720 self._log_prefix, 

721 self.token_label, 

722 ) 

723 await self._safe_refresh_token() 

724 

725 except asyncio.CancelledError: 

726 verbose_proxy_logger.info( 

727 "%s%s refresh loop cancelled", 

728 self._log_prefix, 

729 self.token_label, 

730 ) 

731 break 

732 except Exception as e: 

733 verbose_proxy_logger.error( 

734 "%sError in %s refresh loop: %s. Retrying in %ss...", 

735 self._log_prefix, 

736 self.token_label, 

737 e, 

738 self.FALLBACK_REFRESH_INTERVAL_SECONDS, 

739 ) 

740 # On error, wait before retrying to avoid tight error loops 

741 try: 

742 await asyncio.sleep(self.FALLBACK_REFRESH_INTERVAL_SECONDS) 

743 except asyncio.CancelledError: 

744 break 

745 

746 async def _safe_refresh_token(self) -> None: 

747 """ 

748 Refresh the database token with proper locking to prevent race conditions. 

749 

750 Uses an asyncio lock to ensure only one refresh operation happens at a time, 

751 preventing multiple concurrent reconnection attempts. 

752 """ 

753 async with self._reconnection_lock: 

754 # Double-checked under the lock: another trigger (e.g. the 

755 # proactive loop racing a __getattr__ fallback) may have already 

756 # refreshed while we waited. Recreating again would needlessly kill 

757 # the engine that refresh just spawned (issue #29176), so coalesce 

758 # by skipping when the current token still has comfortable runway. 

759 if self._token_refresh_not_needed(os.getenv(self._db_url_env_var)): 

760 verbose_proxy_logger.debug( 

761 "%s%s still fresh; skipping redundant refresh.", 

762 self._log_prefix, 

763 self.token_label, 

764 ) 

765 return 

766 

767 previous_db_url: Final = os.getenv(self._db_url_env_var) 

768 new_db_url: Final = self.get_rds_iam_token() 

769 if new_db_url: 

770 try: 

771 await self._replace_prisma_client_for_token_refresh_locked(new_db_url) 

772 except (Exception, asyncio.CancelledError): 

773 if previous_db_url is None: 

774 os.environ.pop(self._db_url_env_var, None) 

775 else: 

776 os.environ[self._db_url_env_var] = previous_db_url 

777 raise 

778 self._last_refresh_time = datetime.utcnow() 

779 verbose_proxy_logger.info( 

780 "%s%s refreshed successfully.", 

781 self._log_prefix, 

782 self.token_label, 

783 ) 

784 else: 

785 verbose_proxy_logger.error( 

786 "%sFailed to generate new %s during proactive refresh", 

787 self._log_prefix, 

788 self.token_label, 

789 ) 

790 

791 def _token_refresh_not_needed(self, token_url: str | None) -> bool: 

792 """True iff the token in ``token_url`` has more than the refresh buffer 

793 of runway left, so a refresh would be redundant. 

794 

795 Used to coalesce stacked refresh triggers. Deliberately mirrors the 

796 proactive loop's schedule (refresh at ``expiration - buffer``): a token 

797 with exactly ``buffer`` seconds left is NOT considered fresh, so the 

798 legitimate proactive refresh still fires. Unparseable tokens return 

799 ``False`` (refresh) — skipping them would mean never refreshing. 

800 """ 

801 token: Final = self._extract_token_from_db_url(token_url) 

802 expiration_time: Final = self._parse_token_expiration(token) 

803 if expiration_time is None: 

804 return False 

805 seconds_left: Final = (expiration_time - datetime.utcnow()).total_seconds() 

806 return seconds_left > self.TOKEN_REFRESH_BUFFER_SECONDS 

807 

808 def __getattr__(self, name: str): 

809 """ 

810 Proxy attribute access to the underlying Prisma client. 

811 

812 If IAM token auth is enabled and the token is found expired here, the 

813 proactive refresh task has missed its window. Behavior depends on 

814 whether we're called from inside a running event loop: 

815 

816 - Inside the loop (typical: from a coroutine): schedule a refresh as a 

817 background task and return the (stale) attribute. The caller's await 

818 will likely fail with a connection error and be retried by upper 

819 layers (`call_with_db_reconnect_retry`); by that time the refresh 

820 has either completed or escalated to the proactive loop's error 

821 path. We CANNOT block here — `run_coroutine_threadsafe(...)` + 

822 `future.result()` from inside the same loop deadlocks the loop 

823 (loop thread is blocked, scheduled coroutine never runs, 30s timeout). 

824 

825 - No running loop (sync caller, mostly tests): run the refresh in a 

826 fresh loop and re-fetch the attribute. 

827 """ 

828 original_attr = getattr(self._original_prisma, name) 

829 

830 if self.iam_token_db_auth: 830 ↛ 831line 830 didn't jump to line 831 because the condition on line 830 was never true

831 db_url: Final = os.getenv(self._db_url_env_var) 

832 

833 # Check if token is expired (should be rare if background task is running) 

834 if self.is_token_expired(db_url): 

835 try: 

836 running_loop = asyncio.get_running_loop() 

837 except RuntimeError: 

838 running_loop = None 

839 

840 if running_loop is not None: 

841 verbose_proxy_logger.warning( 

842 "%s%s expired in __getattr__ - proactive refresh " 

843 "may have failed. Scheduling async refresh; the current " 

844 "request may fail and be retried with the fresh token.", 

845 self._log_prefix, 

846 self.token_label, 

847 ) 

848 # Non-blocking: schedule the locked refresh on the 

849 # running loop. The reconnection lock inside 

850 # `_safe_refresh_token` coalesces concurrent triggers. 

851 running_loop.create_task(self._safe_refresh_token()) 

852 else: 

853 verbose_proxy_logger.warning( 

854 "%s%s expired in __getattr__ - proactive refresh " 

855 "may have failed. Triggering synchronous fallback refresh...", 

856 self._log_prefix, 

857 self.token_label, 

858 ) 

859 new_db_url: Final = self.get_rds_iam_token() 

860 if new_db_url: 

861 asyncio.run(self.recreate_prisma_client(new_db_url)) 

862 # Re-fetch attribute against the recreated Prisma instance. 

863 original_attr = getattr(self._original_prisma, name) 

864 verbose_proxy_logger.info( 

865 "%sSynchronous token refresh completed successfully", 

866 self._log_prefix, 

867 ) 

868 else: 

869 raise ValueError(f"Failed to get {self.token_label}") 

870 

871 return original_attr 

872 

873 

874class PrismaManager: 

875 @staticmethod 

876 def _get_prisma_dir() -> str: 

877 """Get the path to the migrations directory""" 

878 abspath: Final = os.path.abspath(__file__) 

879 dname: Final = os.path.dirname(os.path.dirname(abspath)) 

880 return dname 

881 

882 @staticmethod 

883 def _apply_replica_identity_full_if_requested() -> None: 

884 """ 

885 `prisma db push` bypasses litellm-proxy-extras, so the opt-in 

886 REPLICA IDENTITY FULL step has to be driven from here too. 

887 

888 litellm-proxy-extras is an optional install, so this is a no-op when it 

889 is absent. 

890 """ 

891 try: 

892 from litellm_proxy_extras.utils import ProxyExtrasDBManager 

893 except ImportError: 

894 return 

895 ProxyExtrasDBManager.apply_replica_identity_full_if_requested() 

896 

897 @staticmethod 

898 def _raise_if_partitioned_spend_logs() -> None: 

899 """`prisma db push` rewrites a doc-partitioned LiteLLM_SpendLogs 

900 primary key back to ("request_id"), which Postgres rejects. Fail fast 

901 with guidance instead of retrying into that raw error. No-op when 

902 litellm-proxy-extras is absent.""" 

903 try: 

904 from litellm_proxy_extras.utils import ( 

905 PARTITIONED_SPEND_LOGS_PUSH_ERROR, 

906 ProxyExtrasDBManager, 

907 ) 

908 except ImportError: 

909 return 

910 if ProxyExtrasDBManager.spend_logs_is_partitioned(): 

911 raise RuntimeError(PARTITIONED_SPEND_LOGS_PUSH_ERROR) 

912 

913 @staticmethod 

914 def setup_database(use_migrate: bool = False, use_v2_resolver: bool = False) -> bool: 

915 """ 

916 Set up the database using either prisma migrate or prisma db push 

917 

918 Args: 

919 use_migrate: Use `prisma migrate deploy` instead of `db push`. 

920 use_v2_resolver: Opt into the v2 migration resolver that avoids 

921 the diff-and-force recovery behavior (which caused schema 

922 thrashing during rolling deploys). Defaults to False. 

923 

924 Returns: 

925 bool: True if setup was successful, False otherwise 

926 """ 

927 

928 for attempt in range(4): 928 ↛ 982line 928 didn't jump to line 982 because the loop on line 928 didn't complete

929 original_dir = os.getcwd() 

930 prisma_dir = PrismaManager._get_prisma_dir() 

931 os.chdir(prisma_dir) 

932 try: 

933 if use_migrate: 933 ↛ 947line 933 didn't jump to line 947 because the condition on line 933 was always true

934 try: 

935 from litellm_proxy_extras.utils import ProxyExtrasDBManager 

936 except ImportError as e: 

937 verbose_proxy_logger.error("\x1b[1;31mLiteLLM: Failed to import proxy extras. Got %s\x1b[0m", e) 

938 return False 

939 

940 prisma_dir = PrismaManager._get_prisma_dir() 

941 

942 return ProxyExtrasDBManager.setup_database( 

943 use_migrate=use_migrate, 

944 use_v2_resolver=use_v2_resolver, 

945 ) 

946 else: 

947 try: 

948 from litellm_proxy_extras.prisma_toolchain import ( 

949 prisma_command_timeout, 

950 run_prisma, 

951 ) 

952 except ImportError as e: 

953 verbose_proxy_logger.error("\x1b[1;31mLiteLLM: Failed to import proxy extras. Got %s\x1b[0m", e) 

954 return False 

955 

956 PrismaManager._raise_if_partitioned_spend_logs() 

957 run_prisma( 

958 [ 

959 "prisma", 

960 "db", 

961 "push", 

962 "--accept-data-loss", 

963 "--skip-generate", 

964 ], 

965 timeout=prisma_command_timeout(), 

966 env=os.environ.copy(), 

967 stdout=None, 

968 stderr=None, 

969 ) 

970 PrismaManager._apply_replica_identity_full_if_requested() 

971 return True 

972 except subprocess.TimeoutExpired as e: 

973 verbose_proxy_logger.warning("Attempt %s timed out after %.0fs", attempt + 1, e.timeout) 

974 time.sleep(random.randrange(5, 15)) 

975 except subprocess.CalledProcessError as e: 

976 attempts_left = 3 - attempt 

977 retry_msg = f" Retrying... ({attempts_left} attempts left)" if attempts_left > 0 else "" 

978 verbose_proxy_logger.warning("The process failed to execute. Details: %s.%s", e, retry_msg) 

979 time.sleep(random.randrange(5, 15)) 

980 finally: 

981 os.chdir(original_dir) 

982 return False 

983 

984 

985def should_update_prisma_schema( 

986 disable_updates: bool | str | None = None, 

987) -> bool: 

988 """ 

989 Determines if Prisma Schema updates should be applied during startup. 

990 

991 Args: 

992 disable_updates: Controls whether schema updates are disabled. 

993 Accepts boolean or string ('true'/'false'). Defaults to checking DISABLE_SCHEMA_UPDATE env var. 

994 

995 Returns: 

996 bool: True if schema updates should be applied, False if updates are disabled. 

997 

998 Examples: 

999 >>> should_update_prisma_schema() # Checks DISABLE_SCHEMA_UPDATE env var 

1000 >>> should_update_prisma_schema(True) # Explicitly disable updates 

1001 >>> should_update_prisma_schema("false") # Enable updates using string 

1002 """ 

1003 if disable_updates is None: 1003 ↛ 1006line 1003 didn't jump to line 1006 because the condition on line 1003 was always true

1004 disable_updates = os.getenv("DISABLE_SCHEMA_UPDATE", "false") 

1005 

1006 if isinstance(disable_updates, str): 1006 ↛ 1009line 1006 didn't jump to line 1009 because the condition on line 1006 was always true

1007 disable_updates = str_to_bool(disable_updates) 

1008 

1009 return not bool(disable_updates)