refac
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user