This commit is contained in:
Timothy Jaeryang Baek
2026-05-09 04:36:23 +09:00
parent 8ffc3d746f
commit 1b4cd705d0
+15 -7
View File
@@ -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: