From 1b4cd705d0b9a51a5e3a7851ec012fb3141eb0a9 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 9 May 2026 04:36:23 +0900 Subject: [PATCH] refac --- backend/open_webui/routers/tasks.py | 22 +++++++++++++++------- 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/backend/open_webui/routers/tasks.py b/backend/open_webui/routers/tasks.py index b921f7b3e6..1d209aceea 100644 --- a/backend/open_webui/routers/tasks.py +++ b/backend/open_webui/routers/tasks.py @@ -159,6 +159,7 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -187,7 +188,7 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver else: template = DEFAULT_TITLE_GENERATION_PROMPT_TEMPLATE - content = title_generation_template(template, form_data['messages'], user) + content = await title_generation_template(template, form_data['messages'], user) max_tokens = models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) @@ -236,6 +237,7 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -264,7 +266,7 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge else: template = DEFAULT_FOLLOW_UP_GENERATION_PROMPT_TEMPLATE - content = follow_up_generation_template(template, form_data['messages'], user) + content = await follow_up_generation_template(template, form_data['messages'], user) payload = { 'model': task_model_id, @@ -304,6 +306,7 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -332,7 +335,7 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get else: template = DEFAULT_TAGS_GENERATION_PROMPT_TEMPLATE - content = tags_generation_template(template, form_data['messages'], user) + content = await tags_generation_template(template, form_data['messages'], user) payload = { 'model': task_model_id, @@ -366,6 +369,7 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get async def generate_image_prompt(request: Request, form_data: dict, user=Depends(get_verified_user)): if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -394,7 +398,7 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends( else: template = DEFAULT_IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE - content = image_prompt_generation_template(template, form_data['messages'], user) + content = await image_prompt_generation_template(template, form_data['messages'], user) payload = { 'model': task_model_id, @@ -446,6 +450,7 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -474,7 +479,7 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v else: template = DEFAULT_QUERY_GENERATION_PROMPT_TEMPLATE - content = query_generation_template(template, form_data['messages'], user) + content = await query_generation_template(template, form_data['messages'], user) payload = { 'model': task_model_id, @@ -524,6 +529,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -552,7 +558,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend else: template = DEFAULT_AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE - content = autocomplete_generation_template(template, prompt, messages, type, user) + content = await autocomplete_generation_template(template, prompt, messages, type, user) payload = { 'model': task_model_id, @@ -586,6 +592,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend async def generate_emoji(request: Request, form_data: dict, user=Depends(get_verified_user)): if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: @@ -611,7 +618,7 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver template = DEFAULT_EMOJI_GENERATION_PROMPT_TEMPLATE - content = emoji_generation_template(template, form_data['prompt'], user) + content = await emoji_generation_template(template, form_data['prompt'], user) payload = { 'model': task_model_id, @@ -651,6 +658,7 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver async def generate_moa_response(request: Request, form_data: dict, user=Depends(get_verified_user)): if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): models = { + **request.app.state.MODELS, request.state.model['id']: request.state.model, } else: