Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/database/configurations.py: 55%
240 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 02:04 +0000
1from __future__ import annotations
3import logging
4import sqlite3
5import ssl
6import traceback
7from abc import ABC, abstractmethod
8from asyncio import AbstractEventLoop, get_running_loop
9from collections.abc import AsyncGenerator, Hashable
10from contextlib import AbstractAsyncContextManager, asynccontextmanager
11from contextvars import ContextVar
12from functools import partial
13from typing import Any, Optional
15import sqlalchemy as sa
16from sqlalchemy import AdaptedConnection, event
17from sqlalchemy.dialects.sqlite import aiosqlite
18from sqlalchemy.engine.interfaces import DBAPIConnection
19from sqlalchemy.ext.asyncio import (
20 AsyncConnection,
21 AsyncEngine,
22 AsyncSession,
23 AsyncSessionTransaction,
24 create_async_engine,
25)
26from sqlalchemy.pool import ConnectionPoolEntry
27from typing_extensions import TypeAlias
29from prefect._internal.observability import configure_logfire
30from prefect._internal.plugins.manager import (
31 build_manager,
32 call_async_hook,
33 load_entry_point_plugins,
34)
35from prefect.plugins import HookSpec
36from prefect.settings import (
37 PREFECT_API_DATABASE_CONNECTION_TIMEOUT,
38 PREFECT_API_DATABASE_ECHO,
39 PREFECT_API_DATABASE_TIMEOUT,
40 PREFECT_TESTING_UNIT_TEST_MODE,
41 get_current_settings,
42)
43from prefect.utilities.asyncutils import add_event_loop_shutdown_callback
45logfire: Any | None = configure_logfire()
47SQLITE_BEGIN_MODE: ContextVar[Optional[str]] = ContextVar( # novm
48 "SQLITE_BEGIN_MODE", default=None
49)
51_EngineCacheKey: TypeAlias = tuple[AbstractEventLoop, str, bool, Optional[float]]
52ENGINES: dict[_EngineCacheKey, AsyncEngine] = {}
55class ConnectionTracker:
56 """A test utility which tracks the connections given out by a connection pool, to
57 make it easy to see which connections are currently checked out and open."""
59 all_connections: dict[AdaptedConnection, list[str]]
60 open_connections: dict[AdaptedConnection, list[str]]
61 left_field_closes: dict[AdaptedConnection, list[str]]
62 connects: int
63 closes: int
64 active: bool
66 def __init__(self) -> None:
67 self.active = False
68 self.all_connections = {}
69 self.open_connections = {}
70 self.left_field_closes = {}
71 self.connects = 0
72 self.closes = 0
74 def track_pool(self, pool: sa.pool.Pool) -> None:
75 event.listen(pool, "connect", self.on_connect)
76 event.listen(pool, "close", self.on_close)
77 event.listen(pool, "close_detached", self.on_close_detached)
79 def on_connect(
80 self,
81 adapted_connection: AdaptedConnection,
82 connection_record: ConnectionPoolEntry,
83 ) -> None:
84 self.all_connections[adapted_connection] = traceback.format_stack()
85 self.open_connections[adapted_connection] = traceback.format_stack()
86 self.connects += 1
88 def on_close(
89 self,
90 adapted_connection: AdaptedConnection,
91 connection_record: ConnectionPoolEntry,
92 ) -> None:
93 try:
94 del self.open_connections[adapted_connection]
95 except KeyError:
96 self.left_field_closes[adapted_connection] = traceback.format_stack()
97 self.closes += 1
99 def on_close_detached(
100 self,
101 adapted_connection: AdaptedConnection,
102 ) -> None:
103 try:
104 del self.open_connections[adapted_connection]
105 except KeyError:
106 self.left_field_closes[adapted_connection] = traceback.format_stack()
107 self.closes += 1
109 def clear(self) -> None:
110 self.all_connections.clear()
111 self.open_connections.clear()
112 self.left_field_closes.clear()
113 self.connects = 0
114 self.closes = 0
117TRACKER: ConnectionTracker = ConnectionTracker()
120class BaseDatabaseConfiguration(ABC):
121 """
122 Abstract base class used to inject database connection configuration into Prefect.
124 This configuration is responsible for defining how Prefect REST API creates and manages
125 database connections and sessions.
126 """
128 def __init__(
129 self,
130 connection_url: str,
131 echo: Optional[bool] = None,
132 timeout: Optional[float] = None,
133 connection_timeout: Optional[float] = None,
134 sqlalchemy_pool_size: Optional[int] = None,
135 sqlalchemy_max_overflow: Optional[int] = None,
136 connection_app_name: Optional[str] = None,
137 statement_cache_size: Optional[int] = None,
138 prepared_statement_cache_size: Optional[int] = None,
139 search_path: Optional[str] = None,
140 ) -> None:
141 self.connection_url = connection_url
142 self.echo: bool = echo or PREFECT_API_DATABASE_ECHO.value()
143 self.timeout: Optional[float] = timeout or PREFECT_API_DATABASE_TIMEOUT.value()
144 self.connection_timeout: Optional[float] = (
145 connection_timeout or PREFECT_API_DATABASE_CONNECTION_TIMEOUT.value()
146 )
147 self.sqlalchemy_pool_size: Optional[int] = (
148 sqlalchemy_pool_size
149 or get_current_settings().server.database.sqlalchemy.pool_size
150 )
151 self.sqlalchemy_max_overflow: Optional[int] = (
152 sqlalchemy_max_overflow
153 or get_current_settings().server.database.sqlalchemy.max_overflow
154 )
155 self.connection_app_name: Optional[str] = (
156 connection_app_name
157 or get_current_settings().server.database.sqlalchemy.connect_args.application_name
158 )
159 self.statement_cache_size: Optional[int] = (
160 statement_cache_size
161 or get_current_settings().server.database.sqlalchemy.connect_args.statement_cache_size
162 )
163 self.prepared_statement_cache_size: Optional[int] = (
164 prepared_statement_cache_size
165 or get_current_settings().server.database.sqlalchemy.connect_args.prepared_statement_cache_size
166 )
167 self.search_path: Optional[str] = (
168 search_path
169 or get_current_settings().server.database.sqlalchemy.connect_args.search_path
170 )
172 def unique_key(self) -> tuple[Hashable, ...]:
173 """
174 Returns a key used to determine whether to instantiate a new DB interface.
175 """
176 return (self.__class__, self.connection_url)
178 @abstractmethod
179 async def engine(self) -> AsyncEngine:
180 """Returns a SqlAlchemy engine"""
182 @abstractmethod
183 async def session(self, engine: AsyncEngine) -> AsyncSession:
184 """
185 Retrieves a SQLAlchemy session for an engine.
186 """
188 @abstractmethod
189 async def create_db(
190 self, connection: AsyncConnection, base_metadata: sa.MetaData
191 ) -> None:
192 """Create the database"""
194 @abstractmethod
195 async def drop_db(
196 self, connection: AsyncConnection, base_metadata: sa.MetaData
197 ) -> None:
198 """Drop the database"""
200 @abstractmethod
201 def is_inmemory(self) -> bool:
202 """Returns true if database is run in memory"""
204 @abstractmethod
205 def begin_transaction(
206 self, session: AsyncSession, with_for_update: bool = False
207 ) -> AbstractAsyncContextManager[AsyncSessionTransaction]:
208 """Enter a transaction for a session"""
209 pass
212class AsyncPostgresConfiguration(BaseDatabaseConfiguration):
213 async def engine(self) -> AsyncEngine:
214 """Retrieves an async SQLAlchemy engine.
216 Args:
217 connection_url (str, optional): The database connection string.
218 Defaults to self.connection_url
219 echo (bool, optional): Whether to echo SQL sent
220 to the database. Defaults to self.echo
221 timeout (float, optional): The database statement timeout, in seconds.
222 Defaults to self.timeout
224 Returns:
225 AsyncEngine: a SQLAlchemy engine
226 """
228 loop = get_running_loop()
230 cache_key = (
231 loop,
232 self.connection_url,
233 self.echo,
234 self.timeout,
235 )
236 if cache_key not in ENGINES:
237 kwargs: dict[str, Any] = (
238 get_current_settings().server.database.sqlalchemy.model_dump(
239 mode="json", exclude={"connect_args"}
240 )
241 )
242 connect_args: dict[str, Any] = {}
244 if self.timeout is not None:
245 connect_args["command_timeout"] = self.timeout
247 # In test mode, use a higher connection timeout to handle the heavy
248 # load of parallel test execution (pytest-xdist). Establishing a new
249 # asyncpg connection can occasionally take longer than the 5s default
250 # under CI load, which surfaces as a TimeoutError during fixture
251 # setup. Keep the configured value if the user has already raised it.
252 connection_timeout = self.connection_timeout
253 if PREFECT_TESTING_UNIT_TEST_MODE.value() is True: 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true
254 connection_timeout = max(connection_timeout or 0.0, 30.0)
256 if connection_timeout is not None: 256 ↛ 259line 256 didn't jump to line 259 because the condition on line 256 was always true
257 connect_args["timeout"] = connection_timeout
259 if self.statement_cache_size is not None: 259 ↛ 260line 259 didn't jump to line 260 because the condition on line 259 was never true
260 connect_args["statement_cache_size"] = self.statement_cache_size
262 if self.prepared_statement_cache_size is not None: 262 ↛ 263line 262 didn't jump to line 263 because the condition on line 262 was never true
263 connect_args["prepared_statement_cache_size"] = (
264 self.prepared_statement_cache_size
265 )
267 server_settings: dict[str, str] = {}
268 if self.connection_app_name is not None: 268 ↛ 269line 268 didn't jump to line 269 because the condition on line 268 was never true
269 server_settings["application_name"] = self.connection_app_name
270 if self.search_path is not None: 270 ↛ 271line 270 didn't jump to line 271 because the condition on line 270 was never true
271 server_settings["search_path"] = self.search_path
272 if server_settings: 272 ↛ 273line 272 didn't jump to line 273 because the condition on line 272 was never true
273 connect_args["server_settings"] = server_settings
275 if get_current_settings().server.database.sqlalchemy.connect_args.tls.enabled: 275 ↛ 276line 275 didn't jump to line 276 because the condition on line 275 was never true
276 tls_config = (
277 get_current_settings().server.database.sqlalchemy.connect_args.tls
278 )
280 pg_ctx = ssl.create_default_context(purpose=ssl.Purpose.SERVER_AUTH)
282 if tls_config.ca_file: 282 ↛ 287line 282 didn't jump to line 287 because the condition on line 282 was always true
283 pg_ctx = ssl.create_default_context(
284 purpose=ssl.Purpose.SERVER_AUTH, cafile=tls_config.ca_file
285 )
287 pg_ctx.minimum_version = ssl.TLSVersion.TLSv1_2
289 if tls_config.cert_file and tls_config.key_file:
290 pg_ctx.load_cert_chain(
291 certfile=tls_config.cert_file, keyfile=tls_config.key_file
292 )
294 pg_ctx.check_hostname = tls_config.check_hostname
295 pg_ctx.verify_mode = ssl.CERT_REQUIRED
296 connect_args["ssl"] = pg_ctx
298 # Initialize plugin manager
299 if get_current_settings().plugins.enabled: 299 ↛ 300line 299 didn't jump to line 300 because the condition on line 299 was never true
300 pm = build_manager(HookSpec)
301 load_entry_point_plugins(
302 pm,
303 allow=get_current_settings().plugins.allow,
304 deny=get_current_settings().plugins.deny,
305 logger=logging.getLogger("prefect.plugins"),
306 )
308 # Call set_database_connection_params hook
309 results = await call_async_hook(
310 pm,
311 "set_database_connection_params",
312 connection_url=self.connection_url,
313 settings=get_current_settings(),
314 )
316 for _, params, error in results:
317 if error:
318 # Log error but don't fail, other plugins might succeed
319 logging.getLogger("prefect.server.database").warning(
320 "Plugin failed to set database connection params: %s", error
321 )
322 elif params:
323 connect_args.update(params)
325 if connect_args: 325 ↛ 328line 325 didn't jump to line 328 because the condition on line 325 was always true
326 kwargs["connect_args"] = connect_args
328 if self.sqlalchemy_pool_size is not None: 328 ↛ 331line 328 didn't jump to line 331 because the condition on line 328 was always true
329 kwargs["pool_size"] = self.sqlalchemy_pool_size
331 if self.sqlalchemy_max_overflow is not None: 331 ↛ 334line 331 didn't jump to line 334 because the condition on line 331 was always true
332 kwargs["max_overflow"] = self.sqlalchemy_max_overflow
334 engine = create_async_engine(
335 self.connection_url,
336 echo=self.echo,
337 # "pre-ping" connections upon checkout to ensure they have not been
338 # closed on the server side
339 pool_pre_ping=True,
340 # Use connections in LIFO order to help reduce connections
341 # after spiky load and in general increase the likelihood
342 # that a given connection pulled from the pool will be
343 # usable.
344 pool_use_lifo=True,
345 **kwargs,
346 )
348 if logfire: 348 ↛ 349line 348 didn't jump to line 349 because the condition on line 348 was never true
349 logfire.instrument_sqlalchemy(engine) # pyright: ignore
351 if TRACKER.active: 351 ↛ 352line 351 didn't jump to line 352 because the condition on line 351 was never true
352 TRACKER.track_pool(engine.pool)
354 ENGINES[cache_key] = engine
355 await self.schedule_engine_disposal(cache_key)
356 return ENGINES[cache_key]
358 async def schedule_engine_disposal(self, cache_key: _EngineCacheKey) -> None:
359 """
360 Dispose of an engine once the event loop is closing.
362 See caveats at `add_event_loop_shutdown_callback`.
364 We attempted to lazily clean up old engines when new engines are created, but
365 if the loop the engine is attached to is already closed then the connections
366 cannot be cleaned up properly and warnings are displayed.
368 Engine disposal should only be important when running the application
369 ephemerally. Notably, this is an issue in our tests where many short-lived event
370 loops and engines are created which can consume all of the available database
371 connection slots. Users operating at a scale where connection limits are
372 encountered should be encouraged to use a standalone server.
373 """
375 async def dispose_engine(cache_key: _EngineCacheKey) -> None:
376 engine = ENGINES.pop(cache_key, None)
377 if engine:
378 await engine.dispose()
380 await add_event_loop_shutdown_callback(partial(dispose_engine, cache_key))
382 async def session(self, engine: AsyncEngine) -> AsyncSession:
383 """
384 Retrieves a SQLAlchemy session for an engine.
386 Args:
387 engine: a sqlalchemy engine
388 """
389 return AsyncSession(engine, expire_on_commit=False)
391 @asynccontextmanager
392 async def begin_transaction(
393 self, session: AsyncSession, with_for_update: bool = False
394 ) -> AsyncGenerator[AsyncSessionTransaction, None]:
395 # `with_for_update` is for SQLite only. For Postgres, lock the row on read
396 # for update instead.
397 async with session.begin() as transaction:
398 yield transaction
400 async def create_db(
401 self, connection: AsyncConnection, base_metadata: sa.MetaData
402 ) -> None:
403 """Create the database"""
405 await connection.run_sync(base_metadata.create_all)
407 async def drop_db(
408 self, connection: AsyncConnection, base_metadata: sa.MetaData
409 ) -> None:
410 """Drop the database"""
412 await connection.run_sync(base_metadata.drop_all)
414 def is_inmemory(self) -> bool:
415 """Returns true if database is run in memory"""
417 return False
420class AioSqliteConfiguration(BaseDatabaseConfiguration):
421 MIN_SQLITE_VERSION = (3, 24, 0)
423 async def engine(self) -> AsyncEngine:
424 """Retrieves an async SQLAlchemy engine.
426 Args:
427 connection_url (str, optional): The database connection string.
428 Defaults to self.connection_url
429 echo (bool, optional): Whether to echo SQL sent
430 to the database. Defaults to self.echo
431 timeout (float, optional): The database statement timeout, in seconds.
432 Defaults to self.timeout
434 Returns:
435 AsyncEngine: a SQLAlchemy engine
436 """
438 if sqlite3.sqlite_version_info < self.MIN_SQLITE_VERSION:
439 required = ".".join(str(v) for v in self.MIN_SQLITE_VERSION)
440 raise RuntimeError(
441 f"Prefect requires sqlite >= {required} but we found version "
442 f"{sqlite3.sqlite_version}"
443 )
445 kwargs: dict[str, Any] = dict()
447 loop = get_running_loop()
449 cache_key = (loop, self.connection_url, self.echo, self.timeout)
450 if cache_key not in ENGINES:
451 # apply database timeout
452 # In test mode, use a higher timeout to handle lock contention during
453 # parallel test execution. This should match the PRAGMA busy_timeout.
454 if PREFECT_TESTING_UNIT_TEST_MODE.value() is True:
455 kwargs["connect_args"] = dict(timeout=30.0) # 30s for tests
456 elif self.timeout is not None: 456 ↛ anywhereline 456 didn't jump anywhere: it always raised an exception.
457 kwargs["connect_args"] = dict(timeout=self.timeout)
459 # use `named` paramstyle for sqlite instead of `qmark` in very rare
460 # circumstances, we've seen aiosqlite pass parameters in the wrong
461 # order; by using named parameters we avoid this issue
462 # see https://github.com/PrefectHQ/prefect/pull/6702
463 kwargs["paramstyle"] = "named"
465 # ensure a long-lasting pool is used with in-memory databases
466 # because they disappear when the last connection closes
467 if ":memory:" in self.connection_url:
468 kwargs.update(
469 poolclass=sa.pool.AsyncAdaptedQueuePool,
470 pool_size=1,
471 max_overflow=0,
472 pool_recycle=-1,
473 )
475 engine = create_async_engine(self.connection_url, echo=self.echo, **kwargs)
476 event.listen(engine.sync_engine, "connect", self.setup_sqlite)
477 event.listen(engine.sync_engine, "begin", self.begin_sqlite_stmt)
479 if logfire: 479 ↛ 482line 479 didn't jump to line 482 because the condition on line 479 was always true
480 logfire.instrument_sqlalchemy(engine) # pyright: ignore
482 if TRACKER.active:
483 TRACKER.track_pool(engine.pool)
485 ENGINES[cache_key] = engine
486 await self.schedule_engine_disposal(cache_key)
487 return ENGINES[cache_key]
489 async def schedule_engine_disposal(self, cache_key: _EngineCacheKey) -> None:
490 """
491 Dispose of an engine once the event loop is closing.
493 See caveats at `add_event_loop_shutdown_callback`.
495 We attempted to lazily clean up old engines when new engines are created, but
496 if the loop the engine is attached to is already closed then the connections
497 cannot be cleaned up properly and warnings are displayed.
499 Engine disposal should only be important when running the application
500 ephemerally. Notably, this is an issue in our tests where many short-lived event
501 loops and engines are created which can consume all of the available database
502 connection slots. Users operating at a scale where connection limits are
503 encountered should be encouraged to use a standalone server.
504 """
506 async def dispose_engine(cache_key: _EngineCacheKey) -> None:
507 engine = ENGINES.pop(cache_key, None)
508 if engine:
509 await engine.dispose()
511 await add_event_loop_shutdown_callback(partial(dispose_engine, cache_key))
513 def setup_sqlite(self, conn: DBAPIConnection, record: ConnectionPoolEntry) -> None:
514 """Issue PRAGMA statements to SQLITE on connect. PRAGMAs only last for the
515 duration of the connection. See https://www.sqlite.org/pragma.html for more info.
516 """
517 # workaround sqlite transaction behavior
518 if isinstance(conn, aiosqlite.AsyncAdapt_aiosqlite_connection):
519 self.begin_sqlite_conn(conn)
521 cursor = conn.cursor()
523 # write to a write-ahead-log instead and regularly commit the changes
524 # this allows multiple concurrent readers even during a write transaction
525 # even with the WAL we can get busy errors if we have transactions that:
526 # - t1 reads from a database
527 # - t2 inserts to the database
528 # - t1 tries to insert to the database
529 # this can be resolved by using the IMMEDIATE transaction mode in t1
530 cursor.execute("PRAGMA journal_mode = WAL;")
532 # enable foreign keys
533 cursor.execute("PRAGMA foreign_keys = ON;")
535 # disable legacy alter table behavior as it will cause problems during
536 # migrations when tables are renamed as references would otherwise be retained
537 # in some locations
538 # https://www.sqlite.org/pragma.html#pragma_legacy_alter_table
539 cursor.execute("PRAGMA legacy_alter_table=OFF")
541 # when using the WAL, we do need to sync changes on every write. sqlite
542 # recommends using 'normal' mode which is much faster
543 cursor.execute("PRAGMA synchronous = NORMAL;")
545 # a higher cache size (default of 2000) for more aggressive performance
546 cursor.execute("PRAGMA cache_size = 20000;")
548 # wait for this amount of time while a table is locked
549 # before returning and raising an error
550 # setting the value very high allows for more 'concurrency'
551 # without running into errors, but may result in slow api calls
552 if PREFECT_TESTING_UNIT_TEST_MODE.value() is True:
553 cursor.execute("PRAGMA busy_timeout = 30000;") # 30s
554 else:
555 cursor.execute("PRAGMA busy_timeout = 60000;") # 60s
557 # `PRAGMA temp_store = memory;` moves temporary tables from disk into RAM
558 # this supposedly speeds up reads, but it seems to actually
559 # decrease overall performance, see https://github.com/PrefectHQ/prefect/pull/14812
560 # cursor.execute("PRAGMA temp_store = memory;")
562 cursor.close()
564 def begin_sqlite_conn(
565 self, conn: aiosqlite.AsyncAdapt_aiosqlite_connection
566 ) -> None:
567 # disable pysqlite's emitting of the BEGIN statement entirely.
568 # also stops it from emitting COMMIT before any DDL.
569 # requires `begin_sqlite_stmt`
570 # see https://docs.sqlalchemy.org/en/20/dialects/sqlite.html#serializable-isolation-savepoints-transactional-ddl
571 conn.isolation_level = None
573 def begin_sqlite_stmt(self, conn: sa.Connection) -> None:
574 # emit our own BEGIN
575 # requires `begin_sqlite_conn`
576 # see https://docs.sqlalchemy.org/en/20/dialects/sqlite.html#serializable-isolation-savepoints-transactional-ddl
577 mode = SQLITE_BEGIN_MODE.get()
578 if mode is not None:
579 conn.exec_driver_sql(f"BEGIN {mode}")
581 # Note this is intentionally a no-op if there is no BEGIN MODE set
582 # This allows us to use SQLite's default behavior for reads which do not need
583 # to be wrapped in a long-running transaction
585 @asynccontextmanager
586 async def begin_transaction(
587 self, session: AsyncSession, with_for_update: bool = False
588 ) -> AsyncGenerator[AsyncSessionTransaction, None]:
589 token = SQLITE_BEGIN_MODE.set("IMMEDIATE" if with_for_update else "DEFERRED")
591 try:
592 async with session.begin() as transaction:
593 yield transaction
594 finally:
595 SQLITE_BEGIN_MODE.reset(token)
597 async def session(self, engine: AsyncEngine) -> AsyncSession:
598 """
599 Retrieves a SQLAlchemy session for an engine.
601 Args:
602 engine: a sqlalchemy engine
603 """
604 return AsyncSession(engine, expire_on_commit=False)
606 async def create_db(
607 self, connection: AsyncConnection, base_metadata: sa.MetaData
608 ) -> None:
609 """Create the database"""
611 await connection.run_sync(base_metadata.create_all)
613 async def drop_db(
614 self, connection: AsyncConnection, base_metadata: sa.MetaData
615 ) -> None:
616 """Drop the database"""
618 await connection.run_sync(base_metadata.drop_all)
620 def is_inmemory(self) -> bool:
621 """Returns true if database is run in memory"""
623 return ":memory:" in self.connection_url or "mode=memory" in self.connection_url