Coverage for open_webui/tasks.py: 21%
184 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
1# tasks.py
2import asyncio
3import logging
4from contextlib import suppress
5from uuid import uuid4
7from redis.asyncio import Redis
9from open_webui.env import REDIS_KEY_PREFIX, REDIS_RESPONSE_STREAM_TTL, REDIS_TASK_TTL
10from open_webui.utils.json_codec import JSONCodec, dumps_bytes
12log = logging.getLogger(__name__)
14# A dictionary to keep track of active tasks
15tasks: dict[str, asyncio.Task] = {}
16item_tasks = {}
17response_streams: dict[str, dict] = {}
20REDIS_TASKS_KEY = f'{REDIS_KEY_PREFIX}:tasks'
21REDIS_ITEM_TASKS_KEY = f'{REDIS_KEY_PREFIX}:tasks:item'
22REDIS_RESPONSE_STREAMS_KEY = f'{REDIS_KEY_PREFIX}:tasks:response_streams'
23REDIS_PUBSUB_CHANNEL = f'{REDIS_KEY_PREFIX}:tasks:commands'
24REDIS_PUBSUB_RECONNECT_INTERVAL = 1.0
25REDIS_PUBSUB_MAX_RECONNECT_INTERVAL = 30.0
28async def redis_task_command_listener(app):
29 redis: Redis = app.state.redis
30 reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL
32 while True:
33 pubsub = None
34 try:
35 # RedisCluster can't route a pubsub subscribe until initialize() fills its slot cache.
36 await redis.initialize()
38 pubsub = redis.pubsub()
39 await pubsub.subscribe(REDIS_PUBSUB_CHANNEL)
40 reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL
42 async for message in pubsub.listen():
43 if message['type'] != 'message':
44 continue
45 try:
46 command = JSONCodec.loads(message['data'])
47 if command.get('action') != 'stop':
48 continue
50 local_task = tasks.get(command.get('task_id'))
51 if local_task:
52 local_task.cancel()
53 except Exception as e:
54 log.exception(f'Error handling distributed task command: {e}')
55 log.warning('Redis task command listener stopped. Retrying.')
56 except asyncio.CancelledError:
57 raise
58 except Exception as e:
59 log.exception(f'Redis task command listener failed. Retrying: {e}')
60 finally:
61 if pubsub:
62 with suppress(Exception):
63 await pubsub.aclose()
65 await asyncio.sleep(reconnect_interval)
66 reconnect_interval = min(reconnect_interval * 2, REDIS_PUBSUB_MAX_RECONNECT_INTERVAL)
69async def redis_task_heartbeat(app):
70 redis: Redis = app.state.redis
71 while True:
72 await asyncio.sleep(REDIS_TASK_TTL / 4)
73 try:
74 pipe = redis.pipeline(transaction=False)
75 for task_id in list(tasks):
76 # EXPIRE cannot recreate a task already removed by cleanup.
77 pipe.expire(f'{REDIS_TASKS_KEY}:{task_id}', REDIS_TASK_TTL)
78 await pipe.execute()
79 except Exception:
80 log.exception('Redis task heartbeat failed')
83### ------------------------------
84### REDIS-ENABLED HANDLERS
85### ------------------------------
88async def redis_save_task(redis: Redis, task_id: str, item_id: str | None):
89 pipe = redis.pipeline(transaction=False)
90 pipe.set(f'{REDIS_TASKS_KEY}:{task_id}', '1', ex=REDIS_TASK_TTL or None)
91 pipe.hset(REDIS_TASKS_KEY, task_id, item_id or '')
92 if item_id:
93 pipe.sadd(f'{REDIS_ITEM_TASKS_KEY}:{item_id}', task_id)
94 await pipe.execute()
97async def redis_cleanup_task(redis: Redis, task_id: str, item_id: str | None):
98 pipe = redis.pipeline(transaction=False)
99 pipe.delete(f'{REDIS_TASKS_KEY}:{task_id}')
100 pipe.hdel(REDIS_TASKS_KEY, task_id)
101 pipe.hdel(REDIS_RESPONSE_STREAMS_KEY, task_id)
102 if item_id:
103 pipe.srem(f'{REDIS_ITEM_TASKS_KEY}:{item_id}', task_id)
104 await pipe.execute()
107async def redis_list_tasks(redis: Redis, item_id: str | None = None) -> list[str]:
108 task_ids = list(
109 await redis.smembers(f'{REDIS_ITEM_TASKS_KEY}:{item_id}')
110 if item_id is not None
111 else await redis.hkeys(REDIS_TASKS_KEY)
112 )
113 if not task_ids or REDIS_TASK_TTL == 0:
114 return task_ids
116 pipe = redis.pipeline(transaction=False)
117 for task_id in task_ids:
118 pipe.exists(f'{REDIS_TASKS_KEY}:{task_id}')
120 active = []
121 for task_id, exists in zip(task_ids, await pipe.execute()):
122 if exists:
123 active.append(task_id)
124 else:
125 task_item_id = item_id if item_id is not None else await redis.hget(REDIS_TASKS_KEY, task_id)
126 await redis_cleanup_task(redis, task_id, task_item_id or None)
127 return active
130async def redis_send_command(redis: Redis, command: dict):
131 command_json = dumps_bytes(command)
132 # RedisCluster doesn't expose publish() directly, but the
133 # PUBLISH command broadcasts across all cluster nodes server-side.
134 if hasattr(redis, 'nodes_manager'):
135 await redis.execute_command('PUBLISH', REDIS_PUBSUB_CHANNEL, command_json)
136 else:
137 await redis.publish(REDIS_PUBSUB_CHANNEL, command_json)
140async def cleanup_task(redis, task_id: str, id=None):
141 """
142 Remove a completed or canceled task from the global `tasks` dictionary.
143 """
144 if redis:
145 await redis_cleanup_task(redis, task_id, id)
147 tasks.pop(task_id, None) # Remove the task if it exists
148 response_streams.pop(task_id, None)
150 # If an ID is provided, remove the task from the item_tasks dictionary
151 if id and task_id in item_tasks.get(id, []):
152 item_tasks[id].remove(task_id)
153 if not item_tasks[id]: # If no tasks left for this ID, remove the entry
154 item_tasks.pop(id, None)
157async def create_task(redis, coroutine, id=None, task_id=None):
158 """
159 Create a new asyncio task and add it to the global task dictionary.
160 """
161 task_id = task_id or str(uuid4()) # Generate a unique ID for the task
162 task = asyncio.create_task(coroutine) # Create the task
164 # Add a done callback for cleanup
165 task.add_done_callback(lambda t: asyncio.create_task(cleanup_task(redis, task_id, id)))
166 tasks[task_id] = task
168 # If an ID is provided, associate the task with that ID
169 if id:
170 if item_tasks.get(id):
171 item_tasks[id].append(task_id)
172 else:
173 item_tasks[id] = [task_id]
175 if redis:
176 await redis_save_task(redis, task_id, id)
178 return task_id, task
181async def list_tasks(redis):
182 """
183 List all currently active task IDs.
184 """
185 if redis: 185 ↛ 186line 185 didn't jump to line 186 because the condition on line 185 was never true
186 return await redis_list_tasks(redis)
187 return list(tasks.keys())
190async def list_task_ids_by_item_id(redis, id):
191 """
192 List all tasks associated with a specific ID.
193 """
194 if redis: 194 ↛ 195line 194 didn't jump to line 195 because the condition on line 194 was never true
195 return await redis_list_tasks(redis, id)
196 return list(item_tasks.get(id, []))
199async def save_response_stream(
200 redis,
201 task_id: str | None,
202 chat_id: str | None,
203 message_id: str | None,
204 content: str,
205 output: list,
206):
207 if not task_id or not chat_id or not message_id:
208 return
210 data = {
211 'chat_id': chat_id,
212 'message_id': message_id,
213 'content': content,
214 'output': output,
215 }
217 if redis:
218 await redis.hset(REDIS_RESPONSE_STREAMS_KEY, task_id, dumps_bytes(data))
219 if REDIS_RESPONSE_STREAM_TTL > 0:
220 with suppress(Exception):
221 await redis.hexpire(REDIS_RESPONSE_STREAMS_KEY, REDIS_RESPONSE_STREAM_TTL, task_id)
222 else:
223 response_streams[task_id] = data
226async def get_response_streams_by_chat_id(redis, chat_id: str) -> list[dict]:
227 task_ids = await list_task_ids_by_item_id(redis, chat_id)
228 if not task_ids: 228 ↛ 231line 228 didn't jump to line 231 because the condition on line 228 was always true
229 return []
231 if redis:
232 values = await redis.hmget(REDIS_RESPONSE_STREAMS_KEY, task_ids)
233 streams = []
234 for value in values:
235 if not value:
236 continue
237 try:
238 data = JSONCodec.loads(value)
239 except Exception:
240 continue
241 if data.get('chat_id') == chat_id:
242 streams.append(data)
243 return streams
245 return [
246 stream for task_id in task_ids if (stream := response_streams.get(task_id)) and stream.get('chat_id') == chat_id
247 ]
250async def clear_response_stream(redis, task_id: str | None):
251 if not task_id:
252 return
253 if redis:
254 await redis.hdel(REDIS_RESPONSE_STREAMS_KEY, task_id)
255 else:
256 response_streams.pop(task_id, None)
259async def stop_task(redis, task_id: str):
260 """
261 Cancel a running task and remove it from the global task list.
262 """
263 if redis: 263 ↛ 265line 263 didn't jump to line 265 because the condition on line 263 was never true
264 # Look up the item_id before cleanup so we can remove the set entry too
265 item_id = await redis.hget(REDIS_TASKS_KEY, task_id)
266 # PUBSUB: All instances check if they have this task, and stop if so.
267 await redis_send_command(
268 redis,
269 {
270 'action': 'stop',
271 'task_id': task_id,
272 },
273 )
274 # Always clean Redis directly — hdel/srem are idempotent, safe even
275 # if the done_callback on the owning process also fires cleanup.
276 await redis_cleanup_task(redis, task_id, item_id or None)
277 return {'status': True, 'message': f'Task {task_id} stopped.'}
279 task = tasks.pop(task_id, None)
280 if not task: 280 ↛ 283line 280 didn't jump to line 283 because the condition on line 280 was always true
281 return {'status': False, 'message': f'Task with ID {task_id} not found.'}
283 task.cancel() # Request task cancellation
284 try:
285 await task # Wait for the task to handle the cancellation
286 except asyncio.CancelledError:
287 # Task successfully canceled
288 return {'status': True, 'message': f'Task {task_id} successfully stopped.'}
290 if task.cancelled() or task.done():
291 return {'status': True, 'message': f'Task {task_id} successfully cancelled.'}
293 return {'status': True, 'message': f'Cancellation requested for {task_id}.'}
296async def stop_item_tasks(redis: Redis, item_id: str):
297 """
298 Stop all tasks associated with a specific item ID.
299 """
300 task_ids = await list_task_ids_by_item_id(redis, item_id)
301 if not task_ids: 301 ↛ 305line 301 didn't jump to line 305 because the condition on line 301 was always true
302 return {'status': True, 'message': f'No tasks found for item {item_id}.'}
304 # Cleanup mutates the local task list while cancellation is awaited.
305 for task_id in list(task_ids):
306 # A task that already finished needs no stopping; continue with the rest.
307 await stop_task(redis, task_id)
309 return {'status': True, 'message': f'All tasks for item {item_id} stopped.'}
312async def has_active_tasks(redis, chat_id: str) -> bool:
313 """Check if a chat has any active tasks."""
314 task_ids = await list_task_ids_by_item_id(redis, chat_id)
315 return len(task_ids) > 0