Coverage for open_webui/models/automations.py: 55%

223 statements  

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

1import logging 

2import time 

3from typing import Literal, Optional 

4from uuid import uuid4 

5 

6from open_webui.internal.db import Base, get_async_db_context 

7from open_webui.utils.misc import json_text_variants 

8from pydantic import BaseModel, ConfigDict 

9from sqlalchemy import JSON, BigInteger, Boolean, Column, Index, String, Text, cast, delete, func, or_, select, update 

10from sqlalchemy.ext.asyncio import AsyncSession 

11 

12log = logging.getLogger(__name__) 

13 

14 

15#################### 

16# Automation DB Schema 

17#################### 

18 

19 

20class Automation(Base): 

21 __tablename__ = 'automation' 

22 

23 id = Column(Text, primary_key=True) 

24 user_id = Column(Text, nullable=False) 

25 folder_id = Column(Text, nullable=True) 

26 name = Column(Text, nullable=False) 

27 data = Column(JSON, nullable=False) # {prompt, model_id, rrule} 

28 meta = Column(JSON, nullable=True) 

29 is_active = Column(Boolean, nullable=False, default=True) 

30 last_run_at = Column(BigInteger, nullable=True) 

31 next_run_at = Column(BigInteger, nullable=True) 

32 

33 created_at = Column(BigInteger, nullable=False) 

34 updated_at = Column(BigInteger, nullable=False) 

35 

36 __table_args__ = ( 

37 Index('ix_automation_next_run', 'next_run_at'), 

38 Index('ix_automation_user_folder', 'user_id', 'folder_id'), 

39 ) 

40 

41 

42class AutomationRun(Base): 

43 __tablename__ = 'automation_run' 

44 

45 id = Column(Text, primary_key=True) 

46 automation_id = Column(Text, nullable=False) 

47 chat_id = Column(Text, nullable=True) 

48 status = Column(Text, nullable=False) # success | error 

49 error = Column(Text, nullable=True) 

50 created_at = Column(BigInteger, nullable=False) 

51 

52 __table_args__ = ( 

53 Index('ix_automation_run_automation_id', 'automation_id'), 

54 Index('ix_automation_run_aid_created', 'automation_id', 'created_at'), 

55 ) 

56 

57 

58#################### 

59# Pydantic Models 

60#################### 

61 

62 

63class AutomationTerminalConfig(BaseModel): 

64 server_id: str 

65 cwd: Optional[str] = None 

66 

67 

68class AutomationTarget(BaseModel): 

69 type: Literal['chat', 'channel'] = 'chat' 

70 channel_id: Optional[str] = None 

71 

72 

73class AutomationData(BaseModel): 

74 prompt: str 

75 model_id: str 

76 rrule: str 

77 terminal: Optional[AutomationTerminalConfig] = None 

78 target: Optional[AutomationTarget] = None 

79 

80 

81class AutomationModel(BaseModel): 

82 model_config = ConfigDict(from_attributes=True) 

83 

84 id: str 

85 user_id: str 

86 folder_id: Optional[str] = None 

87 name: str 

88 data: dict 

89 meta: Optional[dict] = None 

90 is_active: bool 

91 last_run_at: Optional[int] = None 

92 next_run_at: Optional[int] = None 

93 

94 created_at: int 

95 updated_at: int 

96 

97 

98class AutomationRunModel(BaseModel): 

99 model_config = ConfigDict(from_attributes=True) 

100 

101 id: str 

102 automation_id: str 

103 chat_id: Optional[str] = None 

104 status: str 

105 error: Optional[str] = None 

106 created_at: int 

107 

108 

109class AutomationForm(BaseModel): 

110 name: str 

111 folder_id: Optional[str] = None 

112 data: AutomationData 

113 meta: Optional[dict] = None 

114 is_active: Optional[bool] = True 

115 

116 

117class AutomationResponse(AutomationModel): 

118 last_run: Optional[AutomationRunModel] = None 

119 next_runs: Optional[list[int]] = None 

120 

121 

122class AutomationListResponse(BaseModel): 

123 items: list[AutomationModel] 

124 total: int 

125 

126 

127#################### 

128# AutomationTable 

129#################### 

130 

131 

132class AutomationTable: 

133 async def insert( 

134 self, 

135 user_id: str, 

136 form: AutomationForm, 

137 next_run_at: int, 

138 db: Optional[AsyncSession] = None, 

139 ) -> AutomationModel: 

140 async with get_async_db_context(db) as db: 

141 now = int(time.time_ns()) 

142 row = Automation( 

143 id=str(uuid4()), 

144 user_id=user_id, 

145 folder_id=form.folder_id, 

146 name=form.name, 

147 data=form.data.model_dump(), 

148 meta=form.meta, 

149 is_active=form.is_active, 

150 next_run_at=next_run_at, 

151 created_at=now, 

152 updated_at=now, 

153 ) 

154 db.add(row) 

155 await db.commit() 

156 return AutomationModel.model_validate(row) 

157 

158 async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int: 

159 async with get_async_db_context(db) as db: 

160 result = await db.execute(select(func.count()).select_from(Automation).filter_by(user_id=user_id)) 

161 return result.scalar() 

162 

163 async def get_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationModel]: 

164 async with get_async_db_context(db) as db: 

165 row = await db.get(Automation, id) 

166 return AutomationModel.model_validate(row) if row else None 

167 

168 async def get_active_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> list[AutomationModel]: 

169 """Get active automations for a user (for calendar RRULE expansion).""" 

170 async with get_async_db_context(db) as db: 

171 result = await db.execute( 

172 select(Automation).filter_by(user_id=user_id, is_active=True).order_by(Automation.created_at.desc()) 

173 ) 

174 return [AutomationModel.model_validate(r) for r in result.scalars().all()] 

175 

176 async def search_automations( 

177 self, 

178 user_id: str, 

179 query: Optional[str] = None, 

180 status: Optional[str] = None, 

181 folder_id: Optional[str] = None, 

182 skip: int = 0, 

183 limit: int = 30, 

184 db: Optional[AsyncSession] = None, 

185 ) -> 'AutomationListResponse': 

186 async with get_async_db_context(db) as db: 

187 stmt = select(Automation).filter_by(user_id=user_id) 

188 

189 if folder_id: 

190 stmt = stmt.filter(Automation.folder_id == folder_id) 

191 

192 if query: 

193 # Search the name column and the prompt inside the JSON data. 

194 data_text = cast(Automation.data, String) 

195 stmt = stmt.filter( 

196 or_( 

197 Automation.name.ilike(f'%{query}%'), 

198 *(data_text.ilike(f'%{variant}%') for variant in json_text_variants(query)), 

199 ) 

200 ) 

201 

202 if status == 'active': 202 ↛ 203line 202 didn't jump to line 203 because the condition on line 202 was never true

203 stmt = stmt.filter(Automation.is_active == True) 

204 elif status == 'paused': 204 ↛ 205line 204 didn't jump to line 205 because the condition on line 204 was never true

205 stmt = stmt.filter(Automation.is_active == False) 

206 

207 stmt = stmt.order_by(Automation.created_at.desc()) 

208 

209 # Get total count 

210 count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) 

211 total = count_result.scalar() 

212 

213 if skip: 

214 stmt = stmt.offset(skip) 

215 if limit: 

216 stmt = stmt.limit(limit) 

217 

218 result = await db.execute(stmt) 

219 rows = result.scalars().all() 

220 return AutomationListResponse( 

221 items=[AutomationModel.model_validate(r) for r in rows], 

222 total=total, 

223 ) 

224 

225 async def update_by_id( 

226 self, 

227 id: str, 

228 form: AutomationForm, 

229 next_run_at: int, 

230 db: Optional[AsyncSession] = None, 

231 ) -> Optional[AutomationModel]: 

232 async with get_async_db_context(db) as db: 

233 row = await db.get(Automation, id) 

234 if not row: 

235 return None 

236 row.name = form.name 

237 row.folder_id = form.folder_id 

238 row.data = form.data.model_dump() 

239 row.meta = form.meta 

240 if form.is_active is not None: 

241 row.is_active = form.is_active 

242 row.next_run_at = next_run_at 

243 row.updated_at = int(time.time_ns()) 

244 await db.commit() 

245 return AutomationModel.model_validate(row) 

246 

247 async def clear_folder_ids( 

248 self, 

249 user_id: str, 

250 folder_ids: list[str], 

251 db: Optional[AsyncSession] = None, 

252 ) -> int: 

253 if not folder_ids: 

254 return 0 

255 async with get_async_db_context(db) as db: 

256 result = await db.execute( 

257 update(Automation) 

258 .where(Automation.user_id == user_id, Automation.folder_id.in_(folder_ids)) 

259 .values(folder_id=None, updated_at=int(time.time_ns())) 

260 ) 

261 await db.commit() 

262 return result.rowcount or 0 

263 

264 async def toggle( 

265 self, 

266 id: str, 

267 next_run_at: Optional[int], 

268 db: Optional[AsyncSession] = None, 

269 ) -> Optional[AutomationModel]: 

270 async with get_async_db_context(db) as db: 

271 row = await db.get(Automation, id) 

272 if not row: 

273 return None 

274 row.is_active = not row.is_active 

275 row.next_run_at = next_run_at if row.is_active else None 

276 row.updated_at = int(time.time_ns()) 

277 await db.commit() 

278 return AutomationModel.model_validate(row) 

279 

280 async def delete(self, id: str, db: Optional[AsyncSession] = None) -> bool: 

281 async with get_async_db_context(db) as db: 

282 row = await db.get(Automation, id) 

283 if not row: 

284 return False 

285 await db.delete(row) 

286 await db.commit() 

287 return True 

288 

289 async def claim_due(self, now_ns: int, limit: int = 10, db: Optional[AsyncSession] = None) -> list[AutomationModel]: 

290 """ 

291 Atomically claim due automations for execution. 

292 

293 Advances next_run_at immediately so the row can never be 

294 double-claimed. On PostgreSQL, uses FOR UPDATE SKIP LOCKED 

295 for zero-contention distributed work claiming. 

296 """ 

297 async with get_async_db_context(db) as db: 

298 stmt = ( 

299 select(Automation) 

300 .where( 

301 Automation.is_active == True, 

302 Automation.next_run_at <= now_ns, 

303 ) 

304 .order_by(Automation.next_run_at) 

305 .limit(limit) 

306 ) 

307 

308 if db.bind.dialect.name == 'postgresql': 308 ↛ 309line 308 didn't jump to line 309 because the condition on line 308 was never true

309 stmt = stmt.with_for_update(skip_locked=True) 

310 

311 result = await db.execute(stmt) 

312 rows = result.scalars().all() 

313 

314 from open_webui.utils.automations import next_run_ns 

315 from open_webui.utils.recurrence import RecurrenceEvaluationTimeout 

316 

317 # Batch-fetch user timezones so rescheduling respects each 

318 # user's local timezone instead of falling back to server time. 

319 user_ids = list({row.user_id for row in rows}) 

320 timezone_by_user_id: dict[str, Optional[str]] = {} 

321 if user_ids: 321 ↛ 322line 321 didn't jump to line 322 because the condition on line 321 was never true

322 from open_webui.models.users import User 

323 

324 tz_result = await db.execute(select(User.id, User.timezone).where(User.id.in_(user_ids))) 

325 timezone_by_user_id = {uid: tz for uid, tz in tz_result.all()} 

326 

327 claimed = [] 

328 for row in rows: 328 ↛ 329line 328 didn't jump to line 329 because the loop on line 328 never started

329 try: 

330 next_run_at = await next_run_ns(row.data.get('rrule', ''), tz=timezone_by_user_id.get(row.user_id)) 

331 except RecurrenceEvaluationTimeout: 

332 log.warning('Skipping automation %s: recurrence evaluation timed out', row.id) 

333 continue 

334 row.last_run_at = now_ns 

335 row.next_run_at = next_run_at 

336 claimed.append(row) 

337 

338 await db.commit() 

339 

340 return [AutomationModel.model_validate(r) for r in claimed] 

341 

342 

343#################### 

344# AutomationRunTable 

345#################### 

346 

347 

348class AutomationRunTable: 

349 async def insert( 

350 self, 

351 automation_id: str, 

352 status: str, 

353 chat_id: Optional[str] = None, 

354 error: Optional[str] = None, 

355 db: Optional[AsyncSession] = None, 

356 ) -> AutomationRunModel: 

357 async with get_async_db_context(db) as db: 

358 row = AutomationRun( 

359 id=str(uuid4()), 

360 automation_id=automation_id, 

361 chat_id=chat_id, 

362 status=status, 

363 error=error, 

364 created_at=int(time.time_ns()), 

365 ) 

366 db.add(row) 

367 await db.commit() 

368 return AutomationRunModel.model_validate(row) 

369 

370 async def get_latest(self, automation_id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationRunModel]: 

371 async with get_async_db_context(db) as db: 

372 result = await db.execute( 

373 select(AutomationRun) 

374 .filter_by(automation_id=automation_id) 

375 .order_by(AutomationRun.created_at.desc()) 

376 .limit(1) 

377 ) 

378 row = result.scalars().first() 

379 return AutomationRunModel.model_validate(row) if row else None 

380 

381 async def get_latest_batch( 

382 self, automation_ids: list[str], db: Optional[AsyncSession] = None 

383 ) -> dict[str, AutomationRunModel]: 

384 """Fetch the latest run for each automation in a single query.""" 

385 if not automation_ids: 

386 return {} 

387 async with get_async_db_context(db) as db: 

388 # Subquery: max created_at per automation_id 

389 subq = ( 

390 select( 

391 AutomationRun.automation_id, 

392 func.max(AutomationRun.created_at).label('max_created'), 

393 ) 

394 .filter(AutomationRun.automation_id.in_(automation_ids)) 

395 .group_by(AutomationRun.automation_id) 

396 .subquery() 

397 ) 

398 result = await db.execute( 

399 select(AutomationRun).join( 

400 subq, 

401 (AutomationRun.automation_id == subq.c.automation_id) 

402 & (AutomationRun.created_at == subq.c.max_created), 

403 ) 

404 ) 

405 rows = result.scalars().all() 

406 return {row.automation_id: AutomationRunModel.model_validate(row) for row in rows} 

407 

408 async def get_by_automation( 

409 self, 

410 automation_id: str, 

411 skip: int = 0, 

412 limit: int = 50, 

413 db: Optional[AsyncSession] = None, 

414 ) -> list[AutomationRunModel]: 

415 async with get_async_db_context(db) as db: 

416 result = await db.execute( 

417 select(AutomationRun) 

418 .filter_by(automation_id=automation_id) 

419 .order_by(AutomationRun.created_at.desc()) 

420 .offset(skip) 

421 .limit(limit) 

422 ) 

423 rows = result.scalars().all() 

424 return [AutomationRunModel.model_validate(r) for r in rows] 

425 

426 async def delete_by_automation(self, automation_id: str, db: Optional[AsyncSession] = None) -> int: 

427 async with get_async_db_context(db) as db: 

428 result = await db.execute(delete(AutomationRun).filter_by(automation_id=automation_id)) 

429 await db.commit() 

430 return result.rowcount 

431 

432 async def get_runs_by_user_range( 

433 self, 

434 user_id: str, 

435 start_ns: int, 

436 end_ns: int, 

437 limit: int = 500, 

438 db: Optional[AsyncSession] = None, 

439 ) -> list[tuple['AutomationRunModel', 'AutomationModel']]: 

440 """Get runs within a date range for a user, joined with parent automation.""" 

441 async with get_async_db_context(db) as db: 

442 result = await db.execute( 

443 select(AutomationRun, Automation) 

444 .join(Automation, Automation.id == AutomationRun.automation_id) 

445 .filter( 

446 Automation.user_id == user_id, 

447 AutomationRun.created_at >= start_ns, 

448 AutomationRun.created_at < end_ns, 

449 ) 

450 .order_by(AutomationRun.created_at.desc()) 

451 .limit(limit) 

452 ) 

453 return [ 

454 (AutomationRunModel.model_validate(run), AutomationModel.model_validate(auto)) 

455 for run, auto in result.all() 

456 ] 

457 

458 

459Automations = AutomationTable() 

460AutomationRuns = AutomationRunTable()