feat: automation
This commit is contained in:
@@ -0,0 +1,338 @@
|
||||
"""
|
||||
Automation utilities.
|
||||
|
||||
RRULE helpers, worker loop, and execution logic.
|
||||
Follows the utils/<feature>.py pattern (cf. utils/channels.py, utils/task.py).
|
||||
|
||||
Environment:
|
||||
AUTOMATION_POLL_INTERVAL – seconds between polls (default: 10)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from dateutil.rrule import rrulestr
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from open_webui.models.automations import Automations, AutomationRuns, AutomationModel
|
||||
from open_webui.models.chats import ChatForm, Chats
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.task import prompt_template
|
||||
from open_webui.internal.db import get_db
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
AUTOMATION_POLL_INTERVAL = int(os.getenv('AUTOMATION_POLL_INTERVAL', '10'))
|
||||
|
||||
|
||||
####################
|
||||
# RRULE Helpers
|
||||
####################
|
||||
|
||||
|
||||
def validate_rrule(s: str) -> None:
|
||||
"""Raise ValueError if the RRULE is malformed or exhausted."""
|
||||
try:
|
||||
rule = rrulestr(s, ignoretz=True)
|
||||
except Exception as e:
|
||||
raise ValueError(f'Invalid RRULE: {e}')
|
||||
if rule.after(datetime.now()) is None:
|
||||
raise ValueError('RRULE has no future occurrences')
|
||||
|
||||
|
||||
def next_run_ns(s: str, tz: str = None) -> Optional[int]:
|
||||
"""Next occurrence as epoch nanoseconds, respecting user timezone."""
|
||||
now = datetime.now(ZoneInfo(tz)) if tz else datetime.now()
|
||||
dt = rrulestr(s, ignoretz=True).after(now.replace(tzinfo=None))
|
||||
if dt is None:
|
||||
return None
|
||||
if tz:
|
||||
dt = dt.replace(tzinfo=ZoneInfo(tz))
|
||||
return int(dt.timestamp() * 1_000_000_000)
|
||||
|
||||
|
||||
def next_n_runs_ns(s: str, n: int = 5, tz: str = None) -> list[int]:
|
||||
"""Compute next N occurrences for UI preview."""
|
||||
rule = rrulestr(s, ignoretz=True)
|
||||
result = []
|
||||
dt = datetime.now()
|
||||
for _ in range(n):
|
||||
dt = rule.after(dt)
|
||||
if not dt:
|
||||
break
|
||||
if tz:
|
||||
dt_tz = dt.replace(tzinfo=ZoneInfo(tz))
|
||||
result.append(int(dt_tz.timestamp() * 1_000_000_000))
|
||||
else:
|
||||
result.append(int(dt.timestamp() * 1_000_000_000))
|
||||
return result
|
||||
|
||||
|
||||
############################
|
||||
# Worker Loop
|
||||
############################
|
||||
|
||||
|
||||
async def automation_worker_loop(app) -> None:
|
||||
"""Poll for due automations, claim, fire-and-forget execute.
|
||||
|
||||
Runs on every instance. Poll interval is configurable via
|
||||
AUTOMATION_POLL_INTERVAL env var (default: 10 seconds).
|
||||
"""
|
||||
log.info(
|
||||
f'Automation worker started (poll interval: {AUTOMATION_POLL_INTERVAL}s)'
|
||||
)
|
||||
while True:
|
||||
try:
|
||||
with get_db() as db:
|
||||
batch = Automations.claim_due(
|
||||
int(time.time_ns()), limit=10, db=db
|
||||
)
|
||||
if batch:
|
||||
log.info(f'Claimed {len(batch)} due automation(s)')
|
||||
for automation in batch:
|
||||
asyncio.create_task(execute_automation(app, automation))
|
||||
except Exception:
|
||||
log.exception('Automation worker error')
|
||||
|
||||
# Jitter to spread load across instances
|
||||
await asyncio.sleep(
|
||||
AUTOMATION_POLL_INTERVAL + random.uniform(0, 2)
|
||||
)
|
||||
|
||||
|
||||
##########################
|
||||
# Execute
|
||||
####################
|
||||
|
||||
|
||||
def _build_request(app) -> Request:
|
||||
"""Build a minimal ASGI Request for chat_completion.
|
||||
|
||||
Mirrors the mock-request pattern used in main.py lifespan
|
||||
(model pre-fetch, tool server init) for consistency.
|
||||
"""
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.0"},
|
||||
"method": "POST",
|
||||
"path": "/api/v1/automations/internal",
|
||||
"query_string": b"",
|
||||
"headers": Headers({}).raw,
|
||||
"client": ("127.0.0.1", 0),
|
||||
"server": ("127.0.0.1", 80),
|
||||
"scheme": "http",
|
||||
"app": app,
|
||||
}
|
||||
request = Request(scope)
|
||||
# Ensure request.state is initialized with required attributes
|
||||
request.state.token = None
|
||||
request.state.enable_api_keys = False
|
||||
return request
|
||||
|
||||
|
||||
def _resolve_model_tool_ids(app, model_id: str) -> list[str]:
|
||||
"""Read model-attached tool_ids from model config.
|
||||
|
||||
The frontend does this in Chat.svelte (model.info.meta.toolIds).
|
||||
The backend never auto-resolves them, so we must do it explicitly.
|
||||
"""
|
||||
models = getattr(app.state, "MODELS", {})
|
||||
model = models.get(model_id, {})
|
||||
tool_ids = model.get("info", {}).get("meta", {}).get("toolIds", [])
|
||||
return list(tool_ids) if tool_ids else []
|
||||
|
||||
|
||||
def _resolve_model_features(app, model_id: str) -> dict:
|
||||
"""Read model default features from model config.
|
||||
|
||||
The frontend does this in Chat.svelte (model.info.meta.defaultFeatureIds
|
||||
+ model.info.meta.capabilities). Enables features like web_search,
|
||||
code_interpreter, image_generation when the model has them as defaults
|
||||
AND the capability is enabled AND the admin has enabled the feature.
|
||||
"""
|
||||
models = getattr(app.state, "MODELS", {})
|
||||
model = models.get(model_id, {})
|
||||
meta = model.get("info", {}).get("meta", {})
|
||||
|
||||
default_feature_ids = meta.get("defaultFeatureIds", [])
|
||||
if not default_feature_ids:
|
||||
return {}
|
||||
|
||||
capabilities = meta.get("capabilities", {})
|
||||
config = app.state.config
|
||||
features = {}
|
||||
|
||||
feature_checks = {
|
||||
"web_search": getattr(config, "ENABLE_WEB_SEARCH", False),
|
||||
"image_generation": getattr(config, "ENABLE_IMAGE_GENERATION", False),
|
||||
"code_interpreter": getattr(config, "ENABLE_CODE_INTERPRETER", False),
|
||||
}
|
||||
|
||||
for feature_id in default_feature_ids:
|
||||
if feature_id in feature_checks:
|
||||
# Feature must be: in defaultFeatureIds + capability enabled + admin enabled
|
||||
if capabilities.get(feature_id) and feature_checks[feature_id]:
|
||||
features[feature_id] = True
|
||||
|
||||
return features
|
||||
|
||||
|
||||
def _resolve_model_filter_ids(app, model_id: str) -> list[str]:
|
||||
"""Read model default filter_ids from model config."""
|
||||
models = getattr(app.state, "MODELS", {})
|
||||
model = models.get(model_id, {})
|
||||
filter_ids = model.get("info", {}).get("meta", {}).get("defaultFilterIds", [])
|
||||
return list(filter_ids) if filter_ids else []
|
||||
|
||||
|
||||
async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
"""Execute an automation through the full chat completion pipeline.
|
||||
|
||||
Creates a real chat, then calls chat_completion exactly like the frontend:
|
||||
session_id + chat_id + message_id → async task → pipeline handles everything
|
||||
(filters, model params, knowledge/RAG, tools, DB saves, webhooks).
|
||||
"""
|
||||
try:
|
||||
user = Users.get_user_by_id(automation.user_id)
|
||||
if not user:
|
||||
_record_run(automation.id, "error", error="User not found")
|
||||
return
|
||||
|
||||
prompt = prompt_template(automation.data["prompt"], user)
|
||||
model_id = automation.data["model_id"]
|
||||
|
||||
# Generate proper UUIDs for messages (same as frontend)
|
||||
user_msg_id = str(uuid4())
|
||||
assistant_msg_id = str(uuid4())
|
||||
|
||||
# Create the chat with user message (same structure as frontend)
|
||||
chat = Chats.insert_new_chat(
|
||||
automation.user_id,
|
||||
ChatForm(
|
||||
chat={
|
||||
"title": f"[Automation] {automation.name}",
|
||||
"models": [model_id],
|
||||
"history": {
|
||||
"currentId": user_msg_id,
|
||||
"messages": {
|
||||
user_msg_id: {
|
||||
"id": user_msg_id,
|
||||
"parentId": None,
|
||||
"role": "user",
|
||||
"content": prompt,
|
||||
"childrenIds": [assistant_msg_id],
|
||||
"timestamp": int(time.time()),
|
||||
"models": [model_id],
|
||||
},
|
||||
assistant_msg_id: {
|
||||
"id": assistant_msg_id,
|
||||
"parentId": user_msg_id,
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"model": model_id,
|
||||
"childrenIds": [],
|
||||
"timestamp": int(time.time()),
|
||||
},
|
||||
},
|
||||
},
|
||||
"messages": [
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
"meta": {"automation_id": automation.id},
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
if not chat:
|
||||
_record_run(automation.id, "error", error="Failed to create chat")
|
||||
return
|
||||
|
||||
# Notify frontend to refresh chat list
|
||||
from open_webui.socket.main import sio
|
||||
|
||||
await sio.emit(
|
||||
"events",
|
||||
{
|
||||
"chat_id": chat.id,
|
||||
"message_id": user_msg_id,
|
||||
"data": {"type": "chat:list"},
|
||||
},
|
||||
room=f"user:{automation.user_id}",
|
||||
)
|
||||
|
||||
# Resolve model defaults (frontend does this, backend doesn't)
|
||||
tool_ids = _resolve_model_tool_ids(app, model_id)
|
||||
features = _resolve_model_features(app, model_id)
|
||||
filter_ids = _resolve_model_filter_ids(app, model_id)
|
||||
|
||||
# Build the same payload the frontend sends to /api/chat/completions
|
||||
form_data = {
|
||||
"model": model_id,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"stream": True,
|
||||
"chat_id": chat.id,
|
||||
"id": assistant_msg_id,
|
||||
"parent_id": user_msg_id,
|
||||
"session_id": f"automation:{automation.id}",
|
||||
"background_tasks": {},
|
||||
}
|
||||
if tool_ids:
|
||||
form_data["tool_ids"] = tool_ids
|
||||
if features:
|
||||
form_data["features"] = features
|
||||
if filter_ids:
|
||||
form_data["filter_ids"] = filter_ids
|
||||
|
||||
# Call chat_completion — returns {'status': True, 'task_id': '...'}
|
||||
# The background task handles everything: LLM, tools, DB, webhooks.
|
||||
from open_webui.main import chat_completion as main_chat_completion
|
||||
|
||||
request = _build_request(app)
|
||||
await main_chat_completion(request, form_data, user=user)
|
||||
|
||||
# Notify user
|
||||
from open_webui.socket.main import sio
|
||||
|
||||
await sio.emit(
|
||||
"automation:result",
|
||||
{
|
||||
"automation_id": automation.id,
|
||||
"name": automation.name,
|
||||
"chat_id": chat.id,
|
||||
"status": "success",
|
||||
},
|
||||
room=f"user:{automation.user_id}",
|
||||
)
|
||||
|
||||
_record_run(automation.id, "success", chat_id=chat.id)
|
||||
|
||||
except Exception as e:
|
||||
log.exception(f"Automation {automation.id} failed")
|
||||
_record_run(automation.id, "error", error=str(e)[:4000])
|
||||
|
||||
|
||||
####################
|
||||
# Internals
|
||||
####################
|
||||
|
||||
|
||||
def _record_run(
|
||||
automation_id: str,
|
||||
status: str,
|
||||
chat_id: str = None,
|
||||
error: str = None,
|
||||
):
|
||||
"""Insert a run record into automation_run."""
|
||||
with get_db() as db:
|
||||
AutomationRuns.insert(
|
||||
automation_id, status, chat_id=chat_id, error=error, db=db
|
||||
)
|
||||
@@ -2750,16 +2750,22 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
def get_event_emitter_and_caller(metadata):
|
||||
event_emitter = None
|
||||
event_caller = None
|
||||
if (
|
||||
'session_id' in metadata
|
||||
and metadata['session_id']
|
||||
and 'chat_id' in metadata
|
||||
and metadata['chat_id']
|
||||
and 'message_id' in metadata
|
||||
and metadata['message_id']
|
||||
):
|
||||
|
||||
# event_emitter only needs user_id + chat_id + message_id.
|
||||
# It broadcasts to user:{user_id} room AND persists to DB,
|
||||
# so it works for backend-initiated calls (automations, API).
|
||||
if metadata.get('chat_id') and metadata.get('message_id'):
|
||||
event_emitter = get_event_emitter(metadata)
|
||||
|
||||
# event_caller needs session_id — it calls back to a specific
|
||||
# websocket session (used by direct tools, pyodide code interpreter).
|
||||
if (
|
||||
metadata.get('session_id')
|
||||
and metadata.get('chat_id')
|
||||
and metadata.get('message_id')
|
||||
):
|
||||
event_caller = get_event_call(metadata)
|
||||
|
||||
return event_emitter, event_caller
|
||||
|
||||
|
||||
@@ -3206,7 +3212,9 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
]
|
||||
|
||||
# Standard streaming response handler
|
||||
if event_emitter and event_caller:
|
||||
# event_caller is optional — only needed for direct (client-side) tools
|
||||
# and pyodide code interpreter. Server-side tools work without it.
|
||||
if event_emitter:
|
||||
task_id = str(uuid4()) # Create a unique task ID.
|
||||
model_id = form_data.get('model', '')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user