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
« 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
4import sqlalchemy as sa
5from sqlalchemy.ext.asyncio import AsyncSession
6from sqlalchemy.sql.elements import ColumnElement
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
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 )
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 )
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 )
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 )
59def denied_slots_after_decay(db: PrefectDBInterface) -> ColumnElement[float]:
60 """
61 Calculate denied_slots after applying decay.
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`.
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()
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 )
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 )
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 )
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 )
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
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
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())
131 session.add(model)
132 await session.flush()
134 await emit_concurrency_limit_v2_created_event(model)
136 return model
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")
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()
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)
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)
173 result = await session.execute(query)
174 return result.scalars().unique().all()
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
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")
194 where = (
195 db.ConcurrencyLimitV2.id == concurrency_limit_id
196 if concurrency_limit_id
197 else db.ConcurrencyLimitV2.name == name
198 )
200 await session.execute(
201 sa.update(db.ConcurrencyLimitV2)
202 .where(where)
203 .values(**concurrency_limit.model_dump(exclude_unset=True))
204 )
206 await session.refresh(current_concurrency_limit)
207 await emit_concurrency_limit_v2_updated_event(current_concurrency_limit)
208 return True
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")
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
227 await emit_concurrency_limit_v2_deleted_event(existing)
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
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())
250 return existing_limits
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)
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)
278 result = await session.execute(query)
279 return result.rowcount == len(concurrency_limit_ids)
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 )
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 )
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 )
326 result = await session.execute(query)
327 return result.rowcount == len(concurrency_limit_ids)
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 )
348 result = await session.execute(query)
349 return result.rowcount == len(concurrency_limit_ids)