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
« 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
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
12log = logging.getLogger(__name__)
15####################
16# Automation DB Schema
17####################
20class Automation(Base):
21 __tablename__ = 'automation'
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)
33 created_at = Column(BigInteger, nullable=False)
34 updated_at = Column(BigInteger, nullable=False)
36 __table_args__ = (
37 Index('ix_automation_next_run', 'next_run_at'),
38 Index('ix_automation_user_folder', 'user_id', 'folder_id'),
39 )
42class AutomationRun(Base):
43 __tablename__ = 'automation_run'
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)
52 __table_args__ = (
53 Index('ix_automation_run_automation_id', 'automation_id'),
54 Index('ix_automation_run_aid_created', 'automation_id', 'created_at'),
55 )
58####################
59# Pydantic Models
60####################
63class AutomationTerminalConfig(BaseModel):
64 server_id: str
65 cwd: Optional[str] = None
68class AutomationTarget(BaseModel):
69 type: Literal['chat', 'channel'] = 'chat'
70 channel_id: Optional[str] = None
73class AutomationData(BaseModel):
74 prompt: str
75 model_id: str
76 rrule: str
77 terminal: Optional[AutomationTerminalConfig] = None
78 target: Optional[AutomationTarget] = None
81class AutomationModel(BaseModel):
82 model_config = ConfigDict(from_attributes=True)
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
94 created_at: int
95 updated_at: int
98class AutomationRunModel(BaseModel):
99 model_config = ConfigDict(from_attributes=True)
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
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
117class AutomationResponse(AutomationModel):
118 last_run: Optional[AutomationRunModel] = None
119 next_runs: Optional[list[int]] = None
122class AutomationListResponse(BaseModel):
123 items: list[AutomationModel]
124 total: int
127####################
128# AutomationTable
129####################
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)
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()
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
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()]
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)
189 if folder_id:
190 stmt = stmt.filter(Automation.folder_id == folder_id)
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 )
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)
207 stmt = stmt.order_by(Automation.created_at.desc())
209 # Get total count
210 count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
211 total = count_result.scalar()
213 if skip:
214 stmt = stmt.offset(skip)
215 if limit:
216 stmt = stmt.limit(limit)
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 )
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)
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
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)
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
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.
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 )
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)
311 result = await db.execute(stmt)
312 rows = result.scalars().all()
314 from open_webui.utils.automations import next_run_ns
315 from open_webui.utils.recurrence import RecurrenceEvaluationTimeout
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
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()}
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)
338 await db.commit()
340 return [AutomationModel.model_validate(r) for r in claimed]
343####################
344# AutomationRunTable
345####################
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)
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
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}
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]
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
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 ]
459Automations = AutomationTable()
460AutomationRuns = AutomationRunTable()