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

1# tasks.py 

2import asyncio 

3import logging 

4from contextlib import suppress 

5from uuid import uuid4 

6 

7from redis.asyncio import Redis 

8 

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 

11 

12log = logging.getLogger(__name__) 

13 

14# A dictionary to keep track of active tasks 

15tasks: dict[str, asyncio.Task] = {} 

16item_tasks = {} 

17response_streams: dict[str, dict] = {} 

18 

19 

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 

26 

27 

28async def redis_task_command_listener(app): 

29 redis: Redis = app.state.redis 

30 reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL 

31 

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() 

37 

38 pubsub = redis.pubsub() 

39 await pubsub.subscribe(REDIS_PUBSUB_CHANNEL) 

40 reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL 

41 

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 

49 

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() 

64 

65 await asyncio.sleep(reconnect_interval) 

66 reconnect_interval = min(reconnect_interval * 2, REDIS_PUBSUB_MAX_RECONNECT_INTERVAL) 

67 

68 

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') 

81 

82 

83### ------------------------------ 

84### REDIS-ENABLED HANDLERS 

85### ------------------------------ 

86 

87 

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() 

95 

96 

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() 

105 

106 

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 

115 

116 pipe = redis.pipeline(transaction=False) 

117 for task_id in task_ids: 

118 pipe.exists(f'{REDIS_TASKS_KEY}:{task_id}') 

119 

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 

128 

129 

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) 

138 

139 

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) 

146 

147 tasks.pop(task_id, None) # Remove the task if it exists 

148 response_streams.pop(task_id, None) 

149 

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) 

155 

156 

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 

163 

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 

167 

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] 

174 

175 if redis: 

176 await redis_save_task(redis, task_id, id) 

177 

178 return task_id, task 

179 

180 

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()) 

188 

189 

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, [])) 

197 

198 

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 

209 

210 data = { 

211 'chat_id': chat_id, 

212 'message_id': message_id, 

213 'content': content, 

214 'output': output, 

215 } 

216 

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 

224 

225 

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 [] 

230 

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 

244 

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 ] 

248 

249 

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) 

257 

258 

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.'} 

278 

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.'} 

282 

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.'} 

289 

290 if task.cancelled() or task.done(): 

291 return {'status': True, 'message': f'Task {task_id} successfully cancelled.'} 

292 

293 return {'status': True, 'message': f'Cancellation requested for {task_id}.'} 

294 

295 

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}.'} 

303 

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) 

308 

309 return {'status': True, 'message': f'All tasks for item {item_id} stopped.'} 

310 

311 

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