Coverage for /usr/local/lib/python3.12/site-packages/prefect/server/models/concurrency_limits_v2.py: 77%

103 statements  

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

1from typing import List, Optional, Sequence, Union 

2from uuid import UUID 

3 

4import sqlalchemy as sa 

5from sqlalchemy.ext.asyncio import AsyncSession 

6from sqlalchemy.sql.elements import ColumnElement 

7 

8import prefect.server.schemas as schemas 

9from prefect.server.database import PrefectDBInterface, db_injector, orm_models 

10from prefect.server.events import clients 

11from prefect.server.events.schemas import lifecycle 

12from prefect.server.utilities.database import greatest, least 

13from prefect.settings import get_current_settings 

14from prefect.types._datetime import now 

15 

16 

17async def emit_concurrency_limit_v2_created_event( 

18 concurrency_limit: orm_models.ConcurrencyLimitV2, 

19) -> None: 

20 """Emit an event when a global concurrency limit is created.""" 

21 async with clients.PrefectServerEventsClient() as events_client: 

22 await events_client.emit( 

23 lifecycle.concurrency_limit_v2_created_event(concurrency_limit, now("UTC")) 

24 ) 

25 

26 

27async def emit_concurrency_limit_v2_updated_event( 

28 concurrency_limit: orm_models.ConcurrencyLimitV2, 

29) -> None: 

30 """Emit an event when a global concurrency limit is updated.""" 

31 async with clients.PrefectServerEventsClient() as events_client: 

32 await events_client.emit( 

33 lifecycle.concurrency_limit_v2_updated_event(concurrency_limit, now("UTC")) 

34 ) 

35 

36 

37async def emit_concurrency_limit_v2_deleted_event( 

38 concurrency_limit: orm_models.ConcurrencyLimitV2, 

39) -> None: 

40 """Emit an event when a global concurrency limit is deleted.""" 

41 async with clients.PrefectServerEventsClient() as events_client: 

42 await events_client.emit( 

43 lifecycle.concurrency_limit_v2_deleted_event(concurrency_limit, now("UTC")) 

44 ) 

45 

46 

47def active_slots_after_decay(db: PrefectDBInterface) -> ColumnElement[float]: 

48 # Active slots will decay at a rate of `slot_decay_per_second` per second. 

49 return greatest( 

50 0, 

51 db.ConcurrencyLimitV2.active_slots 

52 - sa.func.floor( 

53 db.ConcurrencyLimitV2.slot_decay_per_second 

54 * sa.func.date_diff_seconds(db.ConcurrencyLimitV2.updated) 

55 ), 

56 ) 

57 

58 

59def denied_slots_after_decay(db: PrefectDBInterface) -> ColumnElement[float]: 

60 """ 

61 Calculate denied_slots after applying decay. 

62 

63 Denied slots decay at a rate of `slot_decay_per_second` per second if it's 

64 greater than 0 (rate limits), otherwise for concurrency limits it decays at 

65 a rate based on clamped `avg_slot_occupancy_seconds`. 

66 

67 The clamping matches the retry-after calculation to prevent denied_slots from 

68 accumulating when clients retry faster than the unclamped decay rate. 

69 """ 

70 settings = get_current_settings() 

71 

72 # Determine max_wait based on limit name prefix 

73 max_wait_for_limit = sa.case( 

74 ( 

75 db.ConcurrencyLimitV2.name.like("tag:%"), 

76 sa.literal(settings.server.tasks.tag_concurrency_slot_wait_seconds), 

77 ), 

78 else_=sa.literal( 

79 settings.server.concurrency.maximum_concurrency_slot_wait_seconds 

80 ), 

81 ) 

82 

83 # Clamp avg_slot_occupancy_seconds with minimum bound to prevent division by zero 

84 clamped_occupancy = greatest( 

85 sa.literal(MINIMUM_OCCUPANCY_SECONDS_PER_SLOT), 

86 least( 

87 sa.cast(db.ConcurrencyLimitV2.avg_slot_occupancy_seconds, sa.Float), 

88 max_wait_for_limit, 

89 ), 

90 ) 

91 

92 # Calculate decay rate: use slot_decay_per_second for rate limits, 

93 # use 1/clamped_occupancy for concurrency limits 

94 decay_rate_per_second = sa.case( 

95 ( 

96 db.ConcurrencyLimitV2.slot_decay_per_second > 0.0, 

97 db.ConcurrencyLimitV2.slot_decay_per_second, # Rate limits - no clamping 

98 ), 

99 else_=(1.0 / clamped_occupancy), # Concurrency limits - use clamped value 

100 ) 

101 

102 return greatest( 

103 0, 

104 db.ConcurrencyLimitV2.denied_slots 

105 - sa.func.floor( 

106 decay_rate_per_second 

107 * sa.func.date_diff_seconds(db.ConcurrencyLimitV2.updated) 

108 ), 

109 ) 

110 

111 

112# OCCUPANCY_SAMPLES_MULTIPLIER is used to determine how many samples to use when 

113# calculating the average occupancy seconds per slot. 

114OCCUPANCY_SAMPLES_MULTIPLIER = 2 

115 

116# MINIMUM_OCCUPANCY_SECONDS_PER_SLOT is used to prevent the average occupancy 

117# from dropping too low and causing divide by zero errors. 

118MINIMUM_OCCUPANCY_SECONDS_PER_SLOT = 0.1 

119 

120 

121@db_injector 

122async def create_concurrency_limit( 

123 db: PrefectDBInterface, 

124 session: AsyncSession, 

125 concurrency_limit: Union[ 

126 schemas.actions.ConcurrencyLimitV2Create, schemas.core.ConcurrencyLimitV2 

127 ], 

128) -> orm_models.ConcurrencyLimitV2: 

129 model = db.ConcurrencyLimitV2(**concurrency_limit.model_dump()) 

130 

131 session.add(model) 

132 await session.flush() 

133 

134 await emit_concurrency_limit_v2_created_event(model) 

135 

136 return model 

137 

138 

139@db_injector 

140async def read_concurrency_limit( 

141 db: PrefectDBInterface, 

142 session: AsyncSession, 

143 concurrency_limit_id: Optional[UUID] = None, 

144 name: Optional[str] = None, 

145) -> Union[orm_models.ConcurrencyLimitV2, None]: 

146 if not concurrency_limit_id and not name: 146 ↛ 147line 146 didn't jump to line 147 because the condition on line 146 was never true

147 raise ValueError("Must provide either concurrency_limit_id or name") 

148 

149 where = ( 

150 db.ConcurrencyLimitV2.id == concurrency_limit_id 

151 if concurrency_limit_id 

152 else db.ConcurrencyLimitV2.name == name 

153 ) 

154 query = sa.select(db.ConcurrencyLimitV2).where(where) 

155 result = await session.execute(query) 

156 return result.scalar() 

157 

158 

159@db_injector 

160async def read_all_concurrency_limits( 

161 db: PrefectDBInterface, 

162 session: AsyncSession, 

163 limit: int, 

164 offset: int, 

165) -> Sequence[orm_models.ConcurrencyLimitV2]: 

166 query = sa.select(db.ConcurrencyLimitV2).order_by(db.ConcurrencyLimitV2.name) 

167 

168 if offset is not None: 168 ↛ 170line 168 didn't jump to line 170 because the condition on line 168 was always true

169 query = query.offset(offset) 

170 if limit is not None: 170 ↛ 173line 170 didn't jump to line 173 because the condition on line 170 was always true

171 query = query.limit(limit) 

172 

173 result = await session.execute(query) 

174 return result.scalars().unique().all() 

175 

176 

177@db_injector 

178async def update_concurrency_limit( 

179 db: PrefectDBInterface, 

180 session: AsyncSession, 

181 concurrency_limit: schemas.actions.ConcurrencyLimitV2Update, 

182 concurrency_limit_id: Optional[UUID] = None, 

183 name: Optional[str] = None, 

184) -> bool: 

185 current_concurrency_limit = await read_concurrency_limit( 

186 session, concurrency_limit_id=concurrency_limit_id, name=name 

187 ) 

188 if not current_concurrency_limit: 188 ↛ 191line 188 didn't jump to line 191 because the condition on line 188 was always true

189 return False 

190 

191 if not concurrency_limit_id and not name: 191 ↛ anywhereline 191 didn't jump anywhere: it always raised an exception.

192 raise ValueError("Must provide either concurrency_limit_id or name") 

193 

194 where = ( 

195 db.ConcurrencyLimitV2.id == concurrency_limit_id 

196 if concurrency_limit_id 

197 else db.ConcurrencyLimitV2.name == name 

198 ) 

199 

200 await session.execute( 

201 sa.update(db.ConcurrencyLimitV2) 

202 .where(where) 

203 .values(**concurrency_limit.model_dump(exclude_unset=True)) 

204 ) 

205 

206 await session.refresh(current_concurrency_limit) 

207 await emit_concurrency_limit_v2_updated_event(current_concurrency_limit) 

208 return True 

209 

210 

211@db_injector 

212async def delete_concurrency_limit( 

213 db: PrefectDBInterface, 

214 session: AsyncSession, 

215 concurrency_limit_id: Optional[UUID] = None, 

216 name: Optional[str] = None, 

217) -> bool: 

218 if not concurrency_limit_id and not name: 218 ↛ 219line 218 didn't jump to line 219 because the condition on line 218 was never true

219 raise ValueError("Must provide either concurrency_limit_id or name") 

220 

221 existing = await read_concurrency_limit( 

222 session, concurrency_limit_id=concurrency_limit_id, name=name 

223 ) 

224 if existing is None: 

225 return False 

226 

227 await emit_concurrency_limit_v2_deleted_event(existing) 

228 

229 where = ( 

230 db.ConcurrencyLimitV2.id == concurrency_limit_id 

231 if concurrency_limit_id 

232 else db.ConcurrencyLimitV2.name == name 

233 ) 

234 await session.execute(sa.delete(db.ConcurrencyLimitV2).where(where)) 

235 return True 

236 

237 

238@db_injector 

239async def bulk_read_concurrency_limits( 

240 db: PrefectDBInterface, 

241 session: AsyncSession, 

242 names: List[str], 

243) -> List[orm_models.ConcurrencyLimitV2]: 

244 # Get all existing concurrency limits in `names`. 

245 existing_query = sa.select(db.ConcurrencyLimitV2).where( 

246 db.ConcurrencyLimitV2.name.in_(names) 

247 ) 

248 existing_limits = list((await session.execute(existing_query)).scalars().all()) 

249 

250 return existing_limits 

251 

252 

253@db_injector 

254async def bulk_increment_active_slots( 

255 db: PrefectDBInterface, 

256 session: AsyncSession, 

257 concurrency_limit_ids: List[UUID], 

258 slots: int, 

259) -> bool: 

260 active_slots = active_slots_after_decay(db) 

261 denied_slots = denied_slots_after_decay(db) 

262 

263 query = ( 

264 sa.update(db.ConcurrencyLimitV2) 

265 .where( 

266 sa.and_( 

267 db.ConcurrencyLimitV2.id.in_(concurrency_limit_ids), 

268 db.ConcurrencyLimitV2.active == True, # noqa 

269 active_slots + slots <= db.ConcurrencyLimitV2.limit, 

270 ) 

271 ) 

272 .values( 

273 active_slots=active_slots + slots, 

274 denied_slots=denied_slots, 

275 ) 

276 ).execution_options(synchronize_session=False) 

277 

278 result = await session.execute(query) 

279 return result.rowcount == len(concurrency_limit_ids) 

280 

281 

282@db_injector 

283async def bulk_decrement_active_slots( 

284 db: PrefectDBInterface, 

285 session: AsyncSession, 

286 concurrency_limit_ids: List[UUID], 

287 slots: int, 

288 occupancy_seconds: Optional[float] = None, 

289) -> bool: 

290 query = ( 

291 sa.update(db.ConcurrencyLimitV2) 

292 .where( 

293 sa.and_( 

294 db.ConcurrencyLimitV2.id.in_(concurrency_limit_ids), 

295 db.ConcurrencyLimitV2.active == True, # noqa 

296 ) 

297 ) 

298 .values( 

299 active_slots=sa.case( 

300 (active_slots_after_decay(db) - slots < 0, 0), 

301 else_=active_slots_after_decay(db) - slots, 

302 ), 

303 denied_slots=denied_slots_after_decay(db), 

304 ) 

305 ) 

306 

307 if occupancy_seconds: 307 ↛ 326line 307 didn't jump to line 326 because the condition on line 307 was always true

308 occupancy_seconds_per_slot = max( 

309 occupancy_seconds / slots, MINIMUM_OCCUPANCY_SECONDS_PER_SLOT 

310 ) 

311 

312 query = query.values( 

313 # Update the average occupancy seconds per slot as a weighted 

314 # average over the last `limit * OCCUPANCY_SAMPLE_MULTIPLIER` samples. 

315 avg_slot_occupancy_seconds=db.ConcurrencyLimitV2.avg_slot_occupancy_seconds 

316 + ( 

317 occupancy_seconds_per_slot 

318 / (db.ConcurrencyLimitV2.limit * OCCUPANCY_SAMPLES_MULTIPLIER) 

319 ) 

320 - ( 

321 db.ConcurrencyLimitV2.avg_slot_occupancy_seconds 

322 / (db.ConcurrencyLimitV2.limit * OCCUPANCY_SAMPLES_MULTIPLIER) 

323 ), 

324 ) 

325 

326 result = await session.execute(query) 

327 return result.rowcount == len(concurrency_limit_ids) 

328 

329 

330@db_injector 

331async def bulk_update_denied_slots( 

332 db: PrefectDBInterface, 

333 session: AsyncSession, 

334 concurrency_limit_ids: List[UUID], 

335 slots: int, 

336) -> bool: 

337 query = ( 

338 sa.update(db.ConcurrencyLimitV2) 

339 .where( 

340 sa.and_( 

341 db.ConcurrencyLimitV2.id.in_(concurrency_limit_ids), 

342 db.ConcurrencyLimitV2.active == True, # noqa 

343 ) 

344 ) 

345 .values(denied_slots=denied_slots_after_decay(db) + slots) 

346 ) 

347 

348 result = await session.execute(query) 

349 return result.rowcount == len(concurrency_limit_ids)