Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/database/interface.py: 87%
176 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 collections.abc import Hashable
2from contextlib import asynccontextmanager
3from typing import TYPE_CHECKING, Any
5import sqlalchemy as sa
6from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession
7from typing_extensions import TypeAlias
9from prefect.server.database import orm_models
10from prefect.server.database.alembic_commands import alembic_downgrade, alembic_upgrade
11from prefect.server.database.configurations import BaseDatabaseConfiguration
12from prefect.server.utilities.database import get_dialect
13from prefect.utilities.asyncutils import run_sync_in_worker_thread
15if TYPE_CHECKING: 15 ↛ 16line 15 didn't jump to line 16 because the condition on line 15 was never true
16 from prefect.server.database.query_components import BaseQueryComponents
18_UniqueKey: TypeAlias = tuple[Hashable, ...]
21class DBSingleton(type):
22 """Ensures that only one database interface is created per unique key"""
24 _instances: dict[tuple[str, _UniqueKey, _UniqueKey, _UniqueKey], "DBSingleton"] = (
25 dict()
26 )
28 def __call__(
29 cls,
30 *args: Any,
31 database_config: BaseDatabaseConfiguration,
32 query_components: "BaseQueryComponents",
33 orm: orm_models.BaseORMConfiguration,
34 **kwargs: Any,
35 ) -> "DBSingleton":
36 instance_key = (
37 cls.__name__,
38 database_config.unique_key(),
39 query_components.unique_key(),
40 orm.unique_key(),
41 )
42 try:
43 instance = cls._instances[instance_key]
44 except KeyError:
45 instance = cls._instances[instance_key] = super().__call__(
46 *args,
47 database_config=database_config,
48 query_components=query_components,
49 orm=orm,
50 **kwargs,
51 )
52 return instance
55class PrefectDBInterface(metaclass=DBSingleton):
56 """
57 An interface for backend-specific SqlAlchemy actions and ORM models.
59 The REST API can be configured to run against different databases in order maintain
60 performance at different scales. This interface integrates database- and dialect-
61 specific configuration into a unified interface that the orchestration engine runs
62 against.
63 """
65 def __init__(
66 self,
67 database_config: BaseDatabaseConfiguration,
68 query_components: "BaseQueryComponents",
69 orm: orm_models.BaseORMConfiguration,
70 ):
71 self.database_config = database_config
72 self.queries = query_components
73 self.orm = orm
75 async def create_db(self) -> None:
76 """Create the database"""
77 await self.run_migrations_upgrade()
79 async def drop_db(self) -> None:
80 """Drop the database by removing all tables directly.
82 This reflects the actual database schema and drops every table rather
83 than running all Alembic downgrade migrations in reverse. Running
84 downgrades is fragile because individual migration downgrade steps may
85 fail on real-world data (e.g. re-adding a foreign key constraint when
86 orphaned references exist). Dropping tables directly is both faster
87 and more robust.
89 Reflection is used instead of `Base.metadata.drop_all()` so that
90 tables created by migrations but not tracked in the ORM (e.g.
91 `deployment_version`, `alembic_version`) are also removed.
92 """
93 engine = await self.engine()
94 async with engine.begin() as conn:
95 # Disable FK checks for SQLite so that tables can be dropped in
96 # any order without triggering constraint errors.
97 dialect = get_dialect(self.database_config.connection_url)
98 is_sqlite = dialect.name == "sqlite"
99 if is_sqlite:
100 await conn.execute(sa.text("PRAGMA foreign_keys = OFF"))
102 try:
103 # Reflect the actual database schema so we capture every
104 # table, including migration-only tables not present in the
105 # ORM metadata.
106 metadata = sa.MetaData()
107 await conn.run_sync(metadata.reflect)
108 await conn.run_sync(metadata.drop_all)
109 finally:
110 if is_sqlite:
111 await conn.execute(sa.text("PRAGMA foreign_keys = ON"))
113 async def run_migrations_upgrade(self) -> None:
114 """Run all upgrade migrations"""
115 await run_sync_in_worker_thread(alembic_upgrade)
117 async def run_migrations_downgrade(self, revision: str = "-1") -> None:
118 """Run all downgrade migrations"""
119 await run_sync_in_worker_thread(alembic_downgrade, revision=revision)
121 async def is_db_connectable(self) -> bool:
122 """
123 Returns boolean indicating if the database is connectable.
124 This method is used to determine if the server is ready to accept requests.
125 """
126 engine = await self.engine()
127 try:
128 async with engine.connect():
129 return True
130 except Exception:
131 return False
133 async def engine(self) -> AsyncEngine:
134 """
135 Provides a SqlAlchemy engine against a specific database.
136 """
137 engine = await self.database_config.engine()
139 return engine
141 async def session(self) -> AsyncSession:
142 """
143 Provides a SQLAlchemy session.
144 """
145 engine = await self.engine()
146 return await self.database_config.session(engine)
148 @asynccontextmanager
149 async def session_context(
150 self, begin_transaction: bool = False, with_for_update: bool = False
151 ):
152 """
153 Provides a SQLAlchemy session and a context manager for opening/closing
154 the underlying connection.
156 Args:
157 begin_transaction: if True, the context manager will begin a SQL transaction.
158 Exiting the context manager will COMMIT or ROLLBACK any changes.
159 """
160 session = await self.session()
161 async with session:
162 if begin_transaction:
163 async with self.database_config.begin_transaction(
164 session, with_for_update=with_for_update
165 ):
166 yield session
167 else:
168 yield session
170 @property
171 def dialect(self) -> type[sa.engine.Dialect]:
172 return get_dialect(self.database_config.connection_url)
174 @property
175 def Base(self) -> type[orm_models.Base]:
176 """Base class for orm models"""
177 return orm_models.Base
179 @property
180 def Flow(self) -> type[orm_models.Flow]:
181 """A flow orm model"""
182 return orm_models.Flow
184 @property
185 def FlowRun(self) -> type[orm_models.FlowRun]:
186 """A flow run orm model"""
187 return orm_models.FlowRun
189 @property
190 def FlowRunState(self) -> type[orm_models.FlowRunState]:
191 """A flow run state orm model"""
192 return orm_models.FlowRunState
194 @property
195 def TaskRun(self) -> type[orm_models.TaskRun]:
196 """A task run orm model"""
197 return orm_models.TaskRun
199 @property
200 def TaskRunState(self) -> type[orm_models.TaskRunState]:
201 """A task run state orm model"""
202 return orm_models.TaskRunState
204 @property
205 def Artifact(self) -> type[orm_models.Artifact]:
206 """An artifact orm model"""
207 return orm_models.Artifact
209 @property
210 def ArtifactCollection(self) -> type[orm_models.ArtifactCollection]:
211 """An artifact collection orm model"""
212 return orm_models.ArtifactCollection
214 @property
215 def TaskRunStateCache(self) -> type[orm_models.TaskRunStateCache]:
216 """A task run state cache orm model"""
217 return orm_models.TaskRunStateCache
219 @property
220 def Deployment(self) -> type[orm_models.Deployment]:
221 """A deployment orm model"""
222 return orm_models.Deployment
224 @property
225 def DeploymentSchedule(self) -> type[orm_models.DeploymentSchedule]:
226 """A deployment schedule orm model"""
227 return orm_models.DeploymentSchedule
229 @property
230 def SavedSearch(self) -> type[orm_models.SavedSearch]:
231 """A saved search orm model"""
232 return orm_models.SavedSearch
234 @property
235 def WorkPool(self) -> type[orm_models.WorkPool]:
236 """A work pool orm model"""
237 return orm_models.WorkPool
239 @property
240 def Worker(self) -> type[orm_models.Worker]:
241 """A worker process orm model"""
242 return orm_models.Worker
244 @property
245 def Log(self) -> type[orm_models.Log]:
246 """A log orm model"""
247 return orm_models.Log
249 @property
250 def ConcurrencyLimit(self) -> type[orm_models.ConcurrencyLimit]:
251 """A concurrency model"""
252 return orm_models.ConcurrencyLimit
254 @property
255 def ConcurrencyLimitV2(self) -> type[orm_models.ConcurrencyLimitV2]:
256 """A v2 concurrency model"""
257 return orm_models.ConcurrencyLimitV2
259 @property
260 def CsrfToken(self) -> type[orm_models.CsrfToken]:
261 """A csrf token model"""
262 return orm_models.CsrfToken
264 @property
265 def WorkQueue(self) -> type[orm_models.WorkQueue]:
266 """A work queue model"""
267 return orm_models.WorkQueue
269 @property
270 def Agent(self) -> type[orm_models.Agent]:
271 """An agent model"""
272 return orm_models.Agent
274 @property
275 def BlockType(self) -> type[orm_models.BlockType]:
276 """A block type model"""
277 return orm_models.BlockType
279 @property
280 def BlockSchema(self) -> type[orm_models.BlockSchema]:
281 """A block schema model"""
282 return orm_models.BlockSchema
284 @property
285 def BlockSchemaReference(self) -> type[orm_models.BlockSchemaReference]:
286 """A block schema reference model"""
287 return orm_models.BlockSchemaReference
289 @property
290 def BlockDocument(self) -> type[orm_models.BlockDocument]:
291 """A block document model"""
292 return orm_models.BlockDocument
294 @property
295 def BlockDocumentReference(self) -> type[orm_models.BlockDocumentReference]:
296 """A block document reference model"""
297 return orm_models.BlockDocumentReference
299 @property
300 def Configuration(self) -> type[orm_models.Configuration]:
301 """An configuration model"""
302 return orm_models.Configuration
304 @property
305 def Variable(self) -> type[orm_models.Variable]:
306 """A variable model"""
307 return orm_models.Variable
309 @property
310 def FlowRunInput(self) -> type[orm_models.FlowRunInput]:
311 """A flow run input model"""
312 return orm_models.FlowRunInput
314 @property
315 def Automation(self) -> type[orm_models.Automation]:
316 """An automation model"""
317 return orm_models.Automation
319 @property
320 def AutomationBucket(self) -> type[orm_models.AutomationBucket]:
321 """An automation bucket model"""
322 return orm_models.AutomationBucket
324 @property
325 def AutomationRelatedResource(self) -> type[orm_models.AutomationRelatedResource]:
326 """An automation related resource model"""
327 return orm_models.AutomationRelatedResource
329 @property
330 def CompositeTriggerChildFiring(
331 self,
332 ) -> type[orm_models.CompositeTriggerChildFiring]:
333 """A model capturing a composite trigger's child firing"""
334 return orm_models.CompositeTriggerChildFiring
336 @property
337 def AutomationEventFollower(self) -> type[orm_models.AutomationEventFollower]:
338 """A model capturing one event following another event"""
339 return orm_models.AutomationEventFollower
341 @property
342 def Event(self) -> type[orm_models.Event]:
343 """An event model"""
344 return orm_models.Event
346 @property
347 def EventResource(self) -> type[orm_models.EventResource]:
348 """An event resource model"""
349 return orm_models.EventResource