From 22052d2c2457403d8bbbda686c68c6f7a978d6e4 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Fri, 11 Jul 2025 18:14:48 +0400 Subject: [PATCH] refac: task --- backend/open_webui/main.py | 4 ++-- backend/open_webui/tasks.py | 48 ++++++++++++++++++------------------- 2 files changed, 26 insertions(+), 26 deletions(-) diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index d95567bb90..c5e980a270 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -453,7 +453,7 @@ from open_webui.utils.redis import get_redis_connection from open_webui.tasks import ( redis_task_command_listener, - list_task_ids_by_chat_id, + list_task_ids_by_item_id, stop_task, list_tasks, ) # Import from tasks.py @@ -1505,7 +1505,7 @@ async def list_tasks_by_chat_id_endpoint( if chat is None or chat.user_id != user.id: return {"task_ids": []} - task_ids = await list_task_ids_by_chat_id(request.app.state.redis, chat_id) + task_ids = await list_task_ids_by_item_id(request.app.state.redis, chat_id) log.debug(f"Task IDs for chat {chat_id}: {task_ids}") return {"task_ids": task_ids} diff --git a/backend/open_webui/tasks.py b/backend/open_webui/tasks.py index 67595f311f..89c4b017ed 100644 --- a/backend/open_webui/tasks.py +++ b/backend/open_webui/tasks.py @@ -16,11 +16,11 @@ log.setLevel(SRC_LOG_LEVELS["MAIN"]) # A dictionary to keep track of active tasks tasks: Dict[str, asyncio.Task] = {} -chat_tasks = {} +item_tasks = {} REDIS_TASKS_KEY = "open-webui:tasks" -REDIS_CHAT_TASKS_KEY = "open-webui:tasks:chat" +REDIS_ITEM_TASKS_KEY = "open-webui:tasks:item" REDIS_PUBSUB_CHANNEL = "open-webui:tasks:commands" @@ -48,21 +48,21 @@ async def redis_task_command_listener(app): ### ------------------------------ -async def redis_save_task(redis: Redis, task_id: str, chat_id: Optional[str]): +async def redis_save_task(redis: Redis, task_id: str, item_id: Optional[str]): pipe = redis.pipeline() - pipe.hset(REDIS_TASKS_KEY, task_id, chat_id or "") - if chat_id: - pipe.sadd(f"{REDIS_CHAT_TASKS_KEY}:{chat_id}", task_id) + pipe.hset(REDIS_TASKS_KEY, task_id, item_id or "") + if item_id: + pipe.sadd(f"{REDIS_ITEM_TASKS_KEY}:{item_id}", task_id) await pipe.execute() -async def redis_cleanup_task(redis: Redis, task_id: str, chat_id: Optional[str]): +async def redis_cleanup_task(redis: Redis, task_id: str, item_id: Optional[str]): pipe = redis.pipeline() pipe.hdel(REDIS_TASKS_KEY, task_id) - if chat_id: - pipe.srem(f"{REDIS_CHAT_TASKS_KEY}:{chat_id}", task_id) - if (await pipe.scard(f"{REDIS_CHAT_TASKS_KEY}:{chat_id}").execute())[-1] == 0: - pipe.delete(f"{REDIS_CHAT_TASKS_KEY}:{chat_id}") # Remove if empty set + if item_id: + pipe.srem(f"{REDIS_ITEM_TASKS_KEY}:{item_id}", task_id) + if (await pipe.scard(f"{REDIS_ITEM_TASKS_KEY}:{item_id}").execute())[-1] == 0: + pipe.delete(f"{REDIS_ITEM_TASKS_KEY}:{item_id}") # Remove if empty set await pipe.execute() @@ -70,8 +70,8 @@ async def redis_list_tasks(redis: Redis) -> List[str]: return list(await redis.hkeys(REDIS_TASKS_KEY)) -async def redis_list_chat_tasks(redis: Redis, chat_id: str) -> List[str]: - return list(await redis.smembers(f"{REDIS_CHAT_TASKS_KEY}:{chat_id}")) +async def redis_list_item_tasks(redis: Redis, item_id: str) -> List[str]: + return list(await redis.smembers(f"{REDIS_ITEM_TASKS_KEY}:{item_id}")) async def redis_send_command(redis: Redis, command: dict): @@ -87,11 +87,11 @@ async def cleanup_task(redis, task_id: str, id=None): tasks.pop(task_id, None) # Remove the task if it exists - # If an ID is provided, remove the task from the chat_tasks dictionary - if id and task_id in chat_tasks.get(id, []): - chat_tasks[id].remove(task_id) - if not chat_tasks[id]: # If no tasks left for this ID, remove the entry - chat_tasks.pop(id, None) + # If an ID is provided, remove the task from the item_tasks dictionary + if id and task_id in item_tasks.get(id, []): + item_tasks[id].remove(task_id) + if not item_tasks[id]: # If no tasks left for this ID, remove the entry + item_tasks.pop(id, None) async def create_task(redis, coroutine, id=None): @@ -108,10 +108,10 @@ async def create_task(redis, coroutine, id=None): tasks[task_id] = task # If an ID is provided, associate the task with that ID - if chat_tasks.get(id): - chat_tasks[id].append(task_id) + if item_tasks.get(id): + item_tasks[id].append(task_id) else: - chat_tasks[id] = [task_id] + item_tasks[id] = [task_id] if redis: await redis_save_task(redis, task_id, id) @@ -128,13 +128,13 @@ async def list_tasks(redis): return list(tasks.keys()) -async def list_task_ids_by_chat_id(redis, id): +async def list_task_ids_by_item_id(redis, id): """ List all tasks associated with a specific ID. """ if redis: - return await redis_list_chat_tasks(redis, id) - return chat_tasks.get(id, []) + return await redis_list_item_tasks(redis, id) + return item_tasks.get(id, []) async def stop_task(redis, task_id: str):