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
« 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"""
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
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
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
34__all__ = (
35 "IAMEndpoint",
36 "PrismaManager",
37 "PrismaWrapper",
38 "parse_iam_endpoint_from_url",
39)
42class _PrismaProcess(Protocol):
43 pid: object
46class _PrismaEngine(Protocol):
47 @property
48 def process(self) -> _PrismaProcess: ... 48 ↛ exitline 48 didn't return from function 'process' because
50 async def query(self, content: str, *, tx_id: str | None) -> object: ... 50 ↛ exitline 50 didn't return from function 'query' because
52 async def start_transaction(self, *, content: str) -> str: ... 52 ↛ exitline 52 didn't return from function 'start_transaction' because
54 async def commit_transaction(self, tx_id: str) -> None: ... 54 ↛ exitline 54 didn't return from function 'commit_transaction' because
56 async def rollback_transaction(self, tx_id: str) -> None: ... 56 ↛ exitline 56 didn't return from function 'rollback_transaction' because
59class _PrismaClient(Protocol):
60 _Prisma__engine: _PrismaEngine
62 @property
63 def _engine(self) -> _PrismaEngine: ... 63 ↛ exitline 63 didn't return from function '_engine' because
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()
73 def begin_operation(self) -> None:
74 if self._active_operations == 0:
75 self._drained.clear()
76 self._active_operations += 1
78 def end_operation(self) -> None:
79 self._active_operations -= 1
80 if self._active_operations == 0:
81 self._drained.set()
83 def transaction_started(self, transaction_id: str) -> None:
84 self._transactions = self._transactions.union((transaction_id,))
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()
92 async def wait_until_drained(self) -> None:
93 await self._drained.wait()
96class _TrackedPrismaEngine:
97 def __init__(self, engine: _PrismaEngine, tracker: _PrismaDrainTracker) -> None:
98 self._engine = engine
99 self.tracker = tracker
101 @property
102 def process(self) -> _PrismaProcess:
103 return self._engine.process
105 def __getattr__(self, name: str) -> object:
106 return getattr(self._engine, name)
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()
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
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)
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)
142class PrismaWrapper:
143 """
144 Wrapper around Prisma client that handles token-based database authentication.
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
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 """
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
159 # Fallback refresh interval if token parsing fails (10 minutes)
160 FALLBACK_REFRESH_INTERVAL_SECONDS = 600
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
167 ENGINE_RETIREMENT_DRAIN_TIMEOUT_SECONDS = 90
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
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 ""
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()
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
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
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"
232 @property
233 def iam_token_db_auth(self) -> bool:
234 """Whether any token strategy is active.
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
241 @staticmethod
242 def _read_engine(prisma_client: _PrismaClient) -> _PrismaEngine:
243 return prisma_client._engine
245 @staticmethod
246 def _write_engine(prisma_client: _PrismaClient, engine: _PrismaEngine) -> None:
247 prisma_client._Prisma__engine = engine
249 def _instrument_prisma_client(self, prisma_client: "Prisma | _PrismaClient") -> _PrismaDrainTracker | None:
250 from prisma.errors import ClientNotConnectedError
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
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.
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
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
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)
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)
305 def _retirement_finished(self, retirement_task: asyncio.Task[None]) -> None:
306 self._retirement_tasks = self._retirement_tasks.difference((retirement_task,))
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)
315 @staticmethod
316 async def _kill_engine_process(pid: int) -> None:
317 """Force-kill the engine subprocess to prevent DB connection pool leaks.
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.
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
349 def _extract_token_from_db_url(self, db_url: str | None) -> str | None:
350 """
351 Extract the token (password) from the DATABASE_URL.
353 The token contains the AWS signature with X-Amz-Date and X-Amz-Expires parameters.
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
370 def _parse_token_expiration(self, token: str | None) -> datetime | None:
371 """
372 Parse the token to extract its expiration time.
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)
380 def _calculate_seconds_until_refresh(self) -> float:
381 """
382 Calculate exactly how many seconds until we need to refresh the token.
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).
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)
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
405 # Calculate when we should refresh (expiration - buffer)
406 refresh_at: Final = expiration_time - timedelta(seconds=self.TOKEN_REFRESH_BUFFER_SECONDS)
408 # How long until refresh time?
409 now: Final = datetime.utcnow()
410 seconds_until_refresh: Final = (refresh_at - now).total_seconds()
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)
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
422 token: Final = self._extract_token_from_db_url(token_url)
423 expiration_time: Final = self._parse_token_expiration(token)
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
430 return datetime.utcnow() > expiration_time
432 def get_rds_iam_token(self) -> str | None:
433 """Mint a fresh database token and update the configured DB URL env var.
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
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
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 )
478 @property
479 def engine_generation(self) -> int:
480 """How many query-engine replacements have completed on this wrapper.
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
489 async def _reconnection_settled(self) -> None:
490 async with self._reconnection_lock:
491 pass
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.
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
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.
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.
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).
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``.
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 )
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`.
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
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
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)
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)
595 await self._original_prisma.connect()
596 self._active_drain_tracker = self._instrument_prisma_client(self._original_prisma)
597 self._engine_generation += 1
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()
605 return True
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
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()
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)
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)
637 self._original_prisma = replacement_prisma
638 self._active_drain_tracker = replacement_drain_tracker
639 self._engine_generation += 1
641 if self.on_engine_replaced is not None:
642 self.on_engine_replaced()
644 self._schedule_engine_retirement(old_engine_pid, old_drain_tracker)
646 async def start_token_refresh_task(self) -> None:
647 """
648 Start the background token refresh task.
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
658 if self._token_refresh_task is not None:
659 verbose_proxy_logger.debug("Token refresh task already running")
660 return
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 )
669 async def stop_token_refresh_task(self) -> None:
670 """
671 Stop the background token refresh task gracefully.
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
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 )
690 async def _token_refresh_loop(self) -> None:
691 """
692 Background loop that proactively refreshes database tokens before expiration.
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 )
705 while True:
706 try:
707 # Calculate exactly how long to sleep until next refresh
708 sleep_seconds = self._calculate_seconds_until_refresh()
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)
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()
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
746 async def _safe_refresh_token(self) -> None:
747 """
748 Refresh the database token with proper locking to prevent race conditions.
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
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 )
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.
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
808 def __getattr__(self, name: str):
809 """
810 Proxy attribute access to the underlying Prisma client.
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:
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).
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)
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)
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
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}")
871 return original_attr
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
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.
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()
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)
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
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.
924 Returns:
925 bool: True if setup was successful, False otherwise
926 """
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
940 prisma_dir = PrismaManager._get_prisma_dir()
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
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
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.
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.
995 Returns:
996 bool: True if schema updates should be applied, False if updates are disabled.
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")
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)
1009 return not bool(disable_updates)