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

89 statements  

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

1""" 

2Functions for interacting with concurrency limit ORM objects. 

3Intended for internal use by the Prefect REST API. 

4""" 

5 

6from datetime import timedelta 

7from typing import List, Optional, Sequence, Union 

8from uuid import UUID 

9 

10import sqlalchemy as sa 

11from sqlalchemy.ext.asyncio import AsyncSession 

12 

13import prefect.server.schemas as schemas 

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

15from prefect.server.events import clients 

16from prefect.server.events.schemas import lifecycle 

17from prefect.types._datetime import now 

18 

19# Clients creating V1 limits can't maintain leases, so we use a long TTL to maintain compatibility. 

20V1_LEASE_TTL = timedelta(days=100 * 365) # ~100 years 

21 

22 

23async def emit_concurrency_limit_created_event( 

24 concurrency_limit: orm_models.ConcurrencyLimit, 

25) -> None: 

26 """Emit an event when a tag-based concurrency limit is created.""" 

27 async with clients.PrefectServerEventsClient() as events_client: 

28 await events_client.emit( 

29 lifecycle.concurrency_limit_created_event(concurrency_limit, now("UTC")) 

30 ) 

31 

32 

33async def emit_concurrency_limit_updated_event( 

34 concurrency_limit: orm_models.ConcurrencyLimit, 

35) -> None: 

36 """Emit an event when a tag-based concurrency limit is updated.""" 

37 async with clients.PrefectServerEventsClient() as events_client: 

38 await events_client.emit( 

39 lifecycle.concurrency_limit_updated_event(concurrency_limit, now("UTC")) 

40 ) 

41 

42 

43async def emit_concurrency_limit_deleted_event( 

44 concurrency_limit: orm_models.ConcurrencyLimit, 

45) -> None: 

46 """Emit an event when a tag-based concurrency limit is deleted.""" 

47 async with clients.PrefectServerEventsClient() as events_client: 

48 await events_client.emit( 

49 lifecycle.concurrency_limit_deleted_event(concurrency_limit, now("UTC")) 

50 ) 

51 

52 

53@db_injector 

54async def create_concurrency_limit( 

55 db: PrefectDBInterface, 

56 session: AsyncSession, 

57 concurrency_limit: schemas.core.ConcurrencyLimit, 

58) -> orm_models.ConcurrencyLimit: 

59 insert_values = concurrency_limit.model_dump_for_orm(exclude_unset=False) 

60 insert_values.pop("created") 

61 insert_values.pop("updated") 

62 concurrency_tag = insert_values["tag"] 

63 

64 # set `updated` manually 

65 # known limitation of `on_conflict_do_update`, will not use `Column.onupdate` 

66 # https://docs.sqlalchemy.org/en/14/dialects/sqlite.html#the-set-clause 

67 concurrency_limit.updated = now("UTC") # type: ignore[assignment] 

68 

69 upsert_start = now("UTC") 

70 insert_stmt = ( 

71 db.queries.insert(db.ConcurrencyLimit) 

72 .values(**insert_values) 

73 .on_conflict_do_update( 

74 index_elements=db.orm.concurrency_limit_unique_upsert_columns, 

75 set_=concurrency_limit.model_dump_for_orm( 

76 include={"concurrency_limit", "updated"} 

77 ), 

78 ) 

79 ) 

80 

81 await session.execute(insert_stmt) 

82 

83 query = ( 

84 sa.select(db.ConcurrencyLimit) 

85 .where(db.ConcurrencyLimit.tag == concurrency_tag) 

86 .execution_options(populate_existing=True) 

87 ) 

88 

89 result = await session.execute(query) 

90 model = result.scalar_one() 

91 

92 if model.created >= upsert_start: 

93 await emit_concurrency_limit_created_event(model) 

94 else: 

95 await emit_concurrency_limit_updated_event(model) 

96 

97 return model 

98 

99 

100@db_injector 

101async def read_concurrency_limit( 

102 db: PrefectDBInterface, 

103 session: AsyncSession, 

104 concurrency_limit_id: UUID, 

105) -> Union[orm_models.ConcurrencyLimit, None]: 

106 """ 

107 Reads a concurrency limit by id. If used for orchestration, simultaneous read race 

108 conditions might allow the concurrency limit to be temporarily exceeded. 

109 """ 

110 

111 query = sa.select(db.ConcurrencyLimit).where( 

112 db.ConcurrencyLimit.id == concurrency_limit_id 

113 ) 

114 

115 result = await session.execute(query) 

116 return result.scalar() 

117 

118 

119@db_injector 

120async def read_concurrency_limit_by_tag( 

121 db: PrefectDBInterface, 

122 session: AsyncSession, 

123 tag: str, 

124) -> Union[orm_models.ConcurrencyLimit, None]: 

125 """ 

126 Reads a concurrency limit by tag. If used for orchestration, simultaneous read race 

127 conditions might allow the concurrency limit to be temporarily exceeded. 

128 """ 

129 

130 query = sa.select(db.ConcurrencyLimit).where(db.ConcurrencyLimit.tag == tag) 

131 

132 result = await session.execute(query) 

133 return result.scalar() 

134 

135 

136@db_injector 

137async def reset_concurrency_limit_by_tag( 

138 db: PrefectDBInterface, 

139 session: AsyncSession, 

140 tag: str, 

141 slot_override: Optional[List[UUID]] = None, 

142) -> Union[orm_models.ConcurrencyLimit, None]: 

143 """ 

144 Resets a concurrency limit by tag. 

145 """ 

146 query = sa.select(db.ConcurrencyLimit).where(db.ConcurrencyLimit.tag == tag) 

147 result = await session.execute(query) 

148 concurrency_limit = result.scalar() 

149 if concurrency_limit: 

150 if slot_override is not None: 

151 concurrency_limit.active_slots = [str(slot) for slot in slot_override] 

152 else: 

153 concurrency_limit.active_slots = [] 

154 return concurrency_limit 

155 

156 

157@db_injector 

158async def filter_concurrency_limits_for_orchestration( 

159 db: PrefectDBInterface, 

160 session: AsyncSession, 

161 tags: List[str], 

162) -> Sequence[orm_models.ConcurrencyLimit]: 

163 """ 

164 Filters concurrency limits by tag. This will apply a "select for update" lock on 

165 these rows to prevent simultaneous read race conditions from enabling the 

166 the concurrency limit on these tags from being temporarily exceeded. 

167 """ 

168 

169 if not tags: 

170 return [] 

171 

172 query = ( 

173 sa.select(db.ConcurrencyLimit) 

174 .filter(db.ConcurrencyLimit.tag.in_(tags)) 

175 .order_by(db.ConcurrencyLimit.tag) 

176 .with_for_update() 

177 ) 

178 result = await session.execute(query) 

179 return result.scalars().all() 

180 

181 

182@db_injector 

183async def delete_concurrency_limit( 

184 db: PrefectDBInterface, 

185 session: AsyncSession, 

186 concurrency_limit_id: UUID, 

187) -> bool: 

188 existing = await read_concurrency_limit( 

189 session, concurrency_limit_id=concurrency_limit_id 

190 ) 

191 if existing is None: 

192 return False 

193 

194 await emit_concurrency_limit_deleted_event(existing) 

195 

196 await session.execute( 

197 sa.delete(db.ConcurrencyLimit).where( 

198 db.ConcurrencyLimit.id == concurrency_limit_id 

199 ) 

200 ) 

201 return True 

202 

203 

204@db_injector 

205async def delete_concurrency_limit_by_tag( 

206 db: PrefectDBInterface, 

207 session: AsyncSession, 

208 tag: str, 

209) -> bool: 

210 existing = await read_concurrency_limit_by_tag(session, tag=tag) 

211 if existing is None: 

212 return False 

213 

214 await emit_concurrency_limit_deleted_event(existing) 

215 

216 await session.execute( 

217 sa.delete(db.ConcurrencyLimit).where(db.ConcurrencyLimit.tag == tag) 

218 ) 

219 return True 

220 

221 

222@db_injector 

223async def read_concurrency_limits( 

224 db: PrefectDBInterface, 

225 session: AsyncSession, 

226 limit: Optional[int] = None, 

227 offset: Optional[int] = None, 

228) -> Sequence[orm_models.ConcurrencyLimit]: 

229 """ 

230 Reads a concurrency limits. If used for orchestration, simultaneous read race 

231 conditions might allow the concurrency limit to be temporarily exceeded. 

232 

233 Args: 

234 session: A database session 

235 offset: Query offset 

236 limit: Query limit 

237 

238 Returns: 

239 List[orm_models.ConcurrencyLimit]: concurrency limits 

240 """ 

241 

242 query = sa.select(db.ConcurrencyLimit).order_by(db.ConcurrencyLimit.tag) 

243 

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

245 query = query.offset(offset) 

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

247 query = query.limit(limit) 

248 

249 result = await session.execute(query) 

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