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

1from collections.abc import Hashable 

2from contextlib import asynccontextmanager 

3from typing import TYPE_CHECKING, Any 

4 

5import sqlalchemy as sa 

6from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession 

7from typing_extensions import TypeAlias 

8 

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 

14 

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 

17 

18_UniqueKey: TypeAlias = tuple[Hashable, ...] 

19 

20 

21class DBSingleton(type): 

22 """Ensures that only one database interface is created per unique key""" 

23 

24 _instances: dict[tuple[str, _UniqueKey, _UniqueKey, _UniqueKey], "DBSingleton"] = ( 

25 dict() 

26 ) 

27 

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 

53 

54 

55class PrefectDBInterface(metaclass=DBSingleton): 

56 """ 

57 An interface for backend-specific SqlAlchemy actions and ORM models. 

58 

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

64 

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 

74 

75 async def create_db(self) -> None: 

76 """Create the database""" 

77 await self.run_migrations_upgrade() 

78 

79 async def drop_db(self) -> None: 

80 """Drop the database by removing all tables directly. 

81 

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. 

88 

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

101 

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

112 

113 async def run_migrations_upgrade(self) -> None: 

114 """Run all upgrade migrations""" 

115 await run_sync_in_worker_thread(alembic_upgrade) 

116 

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) 

120 

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 

132 

133 async def engine(self) -> AsyncEngine: 

134 """ 

135 Provides a SqlAlchemy engine against a specific database. 

136 """ 

137 engine = await self.database_config.engine() 

138 

139 return engine 

140 

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) 

147 

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. 

155 

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 

169 

170 @property 

171 def dialect(self) -> type[sa.engine.Dialect]: 

172 return get_dialect(self.database_config.connection_url) 

173 

174 @property 

175 def Base(self) -> type[orm_models.Base]: 

176 """Base class for orm models""" 

177 return orm_models.Base 

178 

179 @property 

180 def Flow(self) -> type[orm_models.Flow]: 

181 """A flow orm model""" 

182 return orm_models.Flow 

183 

184 @property 

185 def FlowRun(self) -> type[orm_models.FlowRun]: 

186 """A flow run orm model""" 

187 return orm_models.FlowRun 

188 

189 @property 

190 def FlowRunState(self) -> type[orm_models.FlowRunState]: 

191 """A flow run state orm model""" 

192 return orm_models.FlowRunState 

193 

194 @property 

195 def TaskRun(self) -> type[orm_models.TaskRun]: 

196 """A task run orm model""" 

197 return orm_models.TaskRun 

198 

199 @property 

200 def TaskRunState(self) -> type[orm_models.TaskRunState]: 

201 """A task run state orm model""" 

202 return orm_models.TaskRunState 

203 

204 @property 

205 def Artifact(self) -> type[orm_models.Artifact]: 

206 """An artifact orm model""" 

207 return orm_models.Artifact 

208 

209 @property 

210 def ArtifactCollection(self) -> type[orm_models.ArtifactCollection]: 

211 """An artifact collection orm model""" 

212 return orm_models.ArtifactCollection 

213 

214 @property 

215 def TaskRunStateCache(self) -> type[orm_models.TaskRunStateCache]: 

216 """A task run state cache orm model""" 

217 return orm_models.TaskRunStateCache 

218 

219 @property 

220 def Deployment(self) -> type[orm_models.Deployment]: 

221 """A deployment orm model""" 

222 return orm_models.Deployment 

223 

224 @property 

225 def DeploymentSchedule(self) -> type[orm_models.DeploymentSchedule]: 

226 """A deployment schedule orm model""" 

227 return orm_models.DeploymentSchedule 

228 

229 @property 

230 def SavedSearch(self) -> type[orm_models.SavedSearch]: 

231 """A saved search orm model""" 

232 return orm_models.SavedSearch 

233 

234 @property 

235 def WorkPool(self) -> type[orm_models.WorkPool]: 

236 """A work pool orm model""" 

237 return orm_models.WorkPool 

238 

239 @property 

240 def Worker(self) -> type[orm_models.Worker]: 

241 """A worker process orm model""" 

242 return orm_models.Worker 

243 

244 @property 

245 def Log(self) -> type[orm_models.Log]: 

246 """A log orm model""" 

247 return orm_models.Log 

248 

249 @property 

250 def ConcurrencyLimit(self) -> type[orm_models.ConcurrencyLimit]: 

251 """A concurrency model""" 

252 return orm_models.ConcurrencyLimit 

253 

254 @property 

255 def ConcurrencyLimitV2(self) -> type[orm_models.ConcurrencyLimitV2]: 

256 """A v2 concurrency model""" 

257 return orm_models.ConcurrencyLimitV2 

258 

259 @property 

260 def CsrfToken(self) -> type[orm_models.CsrfToken]: 

261 """A csrf token model""" 

262 return orm_models.CsrfToken 

263 

264 @property 

265 def WorkQueue(self) -> type[orm_models.WorkQueue]: 

266 """A work queue model""" 

267 return orm_models.WorkQueue 

268 

269 @property 

270 def Agent(self) -> type[orm_models.Agent]: 

271 """An agent model""" 

272 return orm_models.Agent 

273 

274 @property 

275 def BlockType(self) -> type[orm_models.BlockType]: 

276 """A block type model""" 

277 return orm_models.BlockType 

278 

279 @property 

280 def BlockSchema(self) -> type[orm_models.BlockSchema]: 

281 """A block schema model""" 

282 return orm_models.BlockSchema 

283 

284 @property 

285 def BlockSchemaReference(self) -> type[orm_models.BlockSchemaReference]: 

286 """A block schema reference model""" 

287 return orm_models.BlockSchemaReference 

288 

289 @property 

290 def BlockDocument(self) -> type[orm_models.BlockDocument]: 

291 """A block document model""" 

292 return orm_models.BlockDocument 

293 

294 @property 

295 def BlockDocumentReference(self) -> type[orm_models.BlockDocumentReference]: 

296 """A block document reference model""" 

297 return orm_models.BlockDocumentReference 

298 

299 @property 

300 def Configuration(self) -> type[orm_models.Configuration]: 

301 """An configuration model""" 

302 return orm_models.Configuration 

303 

304 @property 

305 def Variable(self) -> type[orm_models.Variable]: 

306 """A variable model""" 

307 return orm_models.Variable 

308 

309 @property 

310 def FlowRunInput(self) -> type[orm_models.FlowRunInput]: 

311 """A flow run input model""" 

312 return orm_models.FlowRunInput 

313 

314 @property 

315 def Automation(self) -> type[orm_models.Automation]: 

316 """An automation model""" 

317 return orm_models.Automation 

318 

319 @property 

320 def AutomationBucket(self) -> type[orm_models.AutomationBucket]: 

321 """An automation bucket model""" 

322 return orm_models.AutomationBucket 

323 

324 @property 

325 def AutomationRelatedResource(self) -> type[orm_models.AutomationRelatedResource]: 

326 """An automation related resource model""" 

327 return orm_models.AutomationRelatedResource 

328 

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 

335 

336 @property 

337 def AutomationEventFollower(self) -> type[orm_models.AutomationEventFollower]: 

338 """A model capturing one event following another event""" 

339 return orm_models.AutomationEventFollower 

340 

341 @property 

342 def Event(self) -> type[orm_models.Event]: 

343 """An event model""" 

344 return orm_models.Event 

345 

346 @property 

347 def EventResource(self) -> type[orm_models.EventResource]: 

348 """An event resource model""" 

349 return orm_models.EventResource