Coverage for open_webui/routers/tasks.py: 36%

284 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import logging 

2import re 

3from typing import Optional 

4 

5from fastapi import APIRouter, Depends, HTTPException, Request, Response, status 

6from fastapi.responses import JSONResponse, RedirectResponse 

7from open_webui.config import ( 

8 DEFAULT_AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE, 

9 DEFAULT_EMOJI_GENERATION_PROMPT_TEMPLATE, 

10 DEFAULT_FOLLOW_UP_GENERATION_PROMPT_TEMPLATE, 

11 DEFAULT_IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE, 

12 DEFAULT_MOA_GENERATION_PROMPT_TEMPLATE, 

13 DEFAULT_QUERY_GENERATION_PROMPT_TEMPLATE, 

14 DEFAULT_TAGS_GENERATION_PROMPT_TEMPLATE, 

15 DEFAULT_TITLE_GENERATION_PROMPT_TEMPLATE, 

16 DEFAULT_VOICE_MODE_PROMPT_TEMPLATE, 

17) 

18from open_webui.constants import ERROR_MESSAGES, TASKS 

19from open_webui.models.config import Config 

20from open_webui.routers.pipelines import process_pipeline_inlet_filter 

21from open_webui.utils.auth import get_admin_user, get_verified_user 

22from open_webui.utils.chat import generate_chat_completion 

23from open_webui.utils.payload import apply_params_to_form_data 

24from open_webui.utils.task import ( 

25 autocomplete_generation_template, 

26 emoji_generation_template, 

27 follow_up_generation_template, 

28 get_task_model_id, 

29 image_prompt_generation_template, 

30 moa_response_generation_template, 

31 query_generation_template, 

32 tags_generation_template, 

33 title_generation_template, 

34) 

35from pydantic import BaseModel 

36 

37log = logging.getLogger(__name__) 

38 

39router = APIRouter() 

40 

41TASK_CONFIG_KEYS = { 

42 'TASK_MODEL': 'task.model.default', 

43 'TASK_MODEL_EXTERNAL': 'task.model.external', 

44 'TASK_MODEL_PARAMS': 'task.model.params', 

45 'TITLE_GENERATION_PROMPT_TEMPLATE': 'task.title.prompt_template', 

46 'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': 'task.image.prompt_template', 

47 'ENABLE_AUTOCOMPLETE_GENERATION': 'task.autocomplete.enable', 

48 'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': 'task.autocomplete.input_max_length', 

49 'AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE': 'task.autocomplete.prompt_template', 

50 'TAGS_GENERATION_PROMPT_TEMPLATE': 'task.tags.prompt_template', 

51 'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': 'task.follow_up.prompt_template', 

52 'ENABLE_FOLLOW_UP_GENERATION': 'task.follow_up.enable', 

53 'ENABLE_TAGS_GENERATION': 'task.tags.enable', 

54 'ENABLE_TITLE_GENERATION': 'task.title.enable', 

55 'ENABLE_SEARCH_QUERY_GENERATION': 'task.query.search.enable', 

56 'ENABLE_RETRIEVAL_QUERY_GENERATION': 'task.query.retrieval.enable', 

57 'QUERY_GENERATION_PROMPT_TEMPLATE': 'task.query.prompt_template', 

58 'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': 'task.tools.prompt_template', 

59 'ENABLE_VOICE_MODE_PROMPT': 'task.voice.prompt.enable', 

60 'VOICE_MODE_PROMPT_TEMPLATE': 'task.voice.prompt_template', 

61} 

62 

63 

64async def get_config_values(key_map: dict[str, str]) -> dict: 

65 values = await Config.get_many(*key_map.values()) 

66 return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values} 

67 

68 

69def config_updates(data: dict, key_map: dict[str, str]) -> dict: 

70 return {key_map[field]: value for field, value in data.items() if field in key_map} 

71 

72 

73def apply_task_model_params(payload: dict, models: dict, task_model_id: str, params: dict | None = None) -> dict: 

74 model = models.get(payload.get('model')) or models.get(task_model_id) 

75 if not model or (not params and not payload.get('params')): 

76 return payload 

77 return apply_params_to_form_data(payload, model, params or None) 

78 

79 

80async def get_task_model_generation_config(default_model_id: str, models) -> tuple[str, dict]: 

81 config = await Config.get_many( 

82 'task.model.default', 

83 'task.model.external', 

84 'task.model.params', 

85 ) 

86 params = config.get('task.model.params') or {} 

87 if not isinstance(params, dict): 

88 params = {} 

89 

90 return ( 

91 get_task_model_id( 

92 default_model_id, 

93 config.get('task.model.default'), 

94 config.get('task.model.external'), 

95 models, 

96 ), 

97 {key: value for key, value in params.items() if value is not None and value != ''}, 

98 ) 

99 

100 

101################################## 

102# 

103# Task Endpoints 

104# 

105################################## 

106 

107 

108@router.get('/config') 

109async def get_task_config(request: Request, user=Depends(get_verified_user)): 

110 return await get_config_values(TASK_CONFIG_KEYS) 

111 

112 

113class TaskConfigForm(BaseModel): 

114 TASK_MODEL: Optional[str] 

115 TASK_MODEL_EXTERNAL: Optional[str] 

116 TASK_MODEL_PARAMS: dict | None = None 

117 ENABLE_TITLE_GENERATION: bool 

118 TITLE_GENERATION_PROMPT_TEMPLATE: str 

119 IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE: str 

120 ENABLE_AUTOCOMPLETE_GENERATION: bool 

121 AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH: int 

122 AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE: str 

123 TAGS_GENERATION_PROMPT_TEMPLATE: str 

124 FOLLOW_UP_GENERATION_PROMPT_TEMPLATE: str 

125 ENABLE_FOLLOW_UP_GENERATION: bool 

126 ENABLE_TAGS_GENERATION: bool 

127 ENABLE_SEARCH_QUERY_GENERATION: bool 

128 ENABLE_RETRIEVAL_QUERY_GENERATION: bool 

129 QUERY_GENERATION_PROMPT_TEMPLATE: str 

130 TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE: str 

131 ENABLE_VOICE_MODE_PROMPT: bool 

132 VOICE_MODE_PROMPT_TEMPLATE: Optional[str] 

133 

134 

135@router.post('/config/update') 

136async def update_task_config(request: Request, form_data: TaskConfigForm, user=Depends(get_admin_user)): 

137 await Config.upsert(config_updates(form_data.model_dump(), TASK_CONFIG_KEYS)) 

138 return await get_config_values(TASK_CONFIG_KEYS) 

139 

140 

141@router.post('/title/completions') 

142async def generate_title(request: Request, form_data: dict, user=Depends(get_verified_user)): 

143 if not await Config.get('task.title.enable'): 

144 return JSONResponse( 

145 status_code=status.HTTP_200_OK, 

146 content={'detail': 'Title generation is disabled'}, 

147 ) 

148 

149 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 149 ↛ 150line 149 didn't jump to line 150 because the condition on line 149 was never true

150 models = { 

151 **dict(request.app.state.MODELS.items()), 

152 request.state.model['id']: request.state.model, 

153 } 

154 else: 

155 models = request.app.state.MODELS 

156 

157 model_id = form_data['model'] 

158 if not model_id: 

159 raise HTTPException( 

160 status_code=status.HTTP_400_BAD_REQUEST, 

161 detail='No model specified for title generation. Please ensure a model is selected for this chat.', 

162 ) 

163 if model_id not in models: 

164 raise HTTPException( 

165 status_code=status.HTTP_404_NOT_FOUND, 

166 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

167 ) 

168 

169 task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) 

170 

171 log.debug('generating chat title using model %s for user %s ', task_model_id, user.email) 

172 

173 title_template = await Config.get('task.title.prompt_template') 

174 if title_template != '': 

175 template = title_template 

176 else: 

177 template = DEFAULT_TITLE_GENERATION_PROMPT_TEMPLATE 

178 

179 content = await title_generation_template(template, form_data['messages'], user) 

180 task_model_params = task_model_params or { 

181 'max_tokens': models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) 

182 } 

183 

184 payload = { 

185 'model': task_model_id, 

186 'messages': [{'role': 'user', 'content': content}], 

187 'stream': False, 

188 'metadata': { 

189 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

190 'task': str(TASKS.TITLE_GENERATION), 

191 'task_body': form_data, 

192 'chat_id': form_data.get('chat_id', None), 

193 }, 

194 } 

195 

196 # Process the payload through the pipeline 

197 try: 

198 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

199 except Exception as e: 

200 raise e 

201 

202 payload = apply_task_model_params(payload, models, task_model_id, task_model_params) 

203 

204 try: 

205 return await generate_chat_completion(request, form_data=payload, user=user) 

206 except Exception as e: 

207 log.error('Exception occurred', exc_info=True) 

208 return JSONResponse( 

209 status_code=status.HTTP_400_BAD_REQUEST, 

210 content={'detail': 'An internal error has occurred.'}, 

211 ) 

212 

213 

214@router.post('/follow_up/completions') 

215async def generate_follow_ups(request: Request, form_data: dict, user=Depends(get_verified_user)): 

216 if not await Config.get('task.follow_up.enable'): 

217 return JSONResponse( 

218 status_code=status.HTTP_200_OK, 

219 content={'detail': 'Follow-up generation is disabled'}, 

220 ) 

221 

222 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 222 ↛ 223line 222 didn't jump to line 223 because the condition on line 222 was never true

223 models = { 

224 **dict(request.app.state.MODELS.items()), 

225 request.state.model['id']: request.state.model, 

226 } 

227 else: 

228 models = request.app.state.MODELS 

229 

230 model_id = form_data['model'] 

231 if model_id not in models: 

232 raise HTTPException( 

233 status_code=status.HTTP_404_NOT_FOUND, 

234 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

235 ) 

236 

237 task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) 

238 

239 log.debug('generating chat title using model %s for user %s ', task_model_id, user.email) 

240 

241 follow_up_template = await Config.get('task.follow_up.prompt_template') 

242 if follow_up_template != '': 

243 template = follow_up_template 

244 else: 

245 template = DEFAULT_FOLLOW_UP_GENERATION_PROMPT_TEMPLATE 

246 

247 content = await follow_up_generation_template(template, form_data['messages'], user) 

248 

249 payload = { 

250 'model': task_model_id, 

251 'messages': [{'role': 'user', 'content': content}], 

252 'stream': False, 

253 'metadata': { 

254 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

255 'task': str(TASKS.FOLLOW_UP_GENERATION), 

256 'task_body': form_data, 

257 'chat_id': form_data.get('chat_id', None), 

258 }, 

259 } 

260 

261 # Process the payload through the pipeline 

262 try: 

263 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

264 except Exception as e: 

265 raise e 

266 

267 payload = apply_task_model_params(payload, models, task_model_id, task_model_params) 

268 

269 try: 

270 return await generate_chat_completion(request, form_data=payload, user=user) 

271 except Exception as e: 

272 log.error('Exception occurred', exc_info=True) 

273 return JSONResponse( 

274 status_code=status.HTTP_400_BAD_REQUEST, 

275 content={'detail': 'An internal error has occurred.'}, 

276 ) 

277 

278 

279@router.post('/tags/completions') 

280async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get_verified_user)): 

281 if not await Config.get('task.tags.enable'): 

282 return JSONResponse( 

283 status_code=status.HTTP_200_OK, 

284 content={'detail': 'Tags generation is disabled'}, 

285 ) 

286 

287 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 287 ↛ 288line 287 didn't jump to line 288 because the condition on line 287 was never true

288 models = { 

289 **dict(request.app.state.MODELS.items()), 

290 request.state.model['id']: request.state.model, 

291 } 

292 else: 

293 models = request.app.state.MODELS 

294 

295 model_id = form_data['model'] 

296 if model_id not in models: 

297 raise HTTPException( 

298 status_code=status.HTTP_404_NOT_FOUND, 

299 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

300 ) 

301 

302 task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) 

303 

304 log.debug('generating chat tags using model %s for user %s ', task_model_id, user.email) 

305 

306 tags_template = await Config.get('task.tags.prompt_template') 

307 if tags_template != '': 

308 template = tags_template 

309 else: 

310 template = DEFAULT_TAGS_GENERATION_PROMPT_TEMPLATE 

311 

312 content = await tags_generation_template(template, form_data['messages'], user) 

313 

314 payload = { 

315 'model': task_model_id, 

316 'messages': [{'role': 'user', 'content': content}], 

317 'stream': False, 

318 'metadata': { 

319 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

320 'task': str(TASKS.TAGS_GENERATION), 

321 'task_body': form_data, 

322 'chat_id': form_data.get('chat_id', None), 

323 }, 

324 } 

325 

326 # Process the payload through the pipeline 

327 try: 

328 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

329 except Exception as e: 

330 raise e 

331 

332 payload = apply_task_model_params(payload, models, task_model_id, task_model_params) 

333 

334 try: 

335 return await generate_chat_completion(request, form_data=payload, user=user) 

336 except Exception as e: 

337 log.error(f'Error generating chat completion: {e}') 

338 return JSONResponse( 

339 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, 

340 content={'detail': 'An internal error has occurred.'}, 

341 ) 

342 

343 

344@router.post('/image_prompt/completions') 

345async def generate_image_prompt(request: Request, form_data: dict, user=Depends(get_verified_user)): 

346 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 346 ↛ 347line 346 didn't jump to line 347 because the condition on line 346 was never true

347 models = { 

348 **dict(request.app.state.MODELS.items()), 

349 request.state.model['id']: request.state.model, 

350 } 

351 else: 

352 models = request.app.state.MODELS 

353 

354 model_id = form_data['model'] 

355 if model_id not in models: 

356 raise HTTPException( 

357 status_code=status.HTTP_404_NOT_FOUND, 

358 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

359 ) 

360 

361 task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) 

362 

363 log.debug('generating image prompt using model %s for user %s ', task_model_id, user.email) 

364 

365 image_prompt_template = await Config.get('task.image.prompt_template') 

366 if image_prompt_template != '': 

367 template = image_prompt_template 

368 else: 

369 template = DEFAULT_IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE 

370 

371 content = await image_prompt_generation_template(template, form_data['messages'], user) 

372 

373 payload = { 

374 'model': task_model_id, 

375 'messages': [{'role': 'user', 'content': content}], 

376 'stream': False, 

377 'metadata': { 

378 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

379 'task': str(TASKS.IMAGE_PROMPT_GENERATION), 

380 'task_body': form_data, 

381 'chat_id': form_data.get('chat_id', None), 

382 }, 

383 } 

384 

385 # Process the payload through the pipeline 

386 try: 

387 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

388 except Exception as e: 

389 raise e 

390 

391 payload = apply_task_model_params(payload, models, task_model_id, task_model_params) 

392 

393 try: 

394 return await generate_chat_completion(request, form_data=payload, user=user) 

395 except Exception as e: 

396 log.error('Exception occurred', exc_info=True) 

397 return JSONResponse( 

398 status_code=status.HTTP_400_BAD_REQUEST, 

399 content={'detail': 'An internal error has occurred.'}, 

400 ) 

401 

402 

403@router.post('/queries/completions') 

404async def generate_queries(request: Request, form_data: dict, user=Depends(get_verified_user)): 

405 type = form_data.get('type') 

406 if type == 'web_search': 406 ↛ 407line 406 didn't jump to line 407 because the condition on line 406 was never true

407 if not await Config.get('task.query.search.enable'): 

408 raise HTTPException( 

409 status_code=status.HTTP_400_BAD_REQUEST, 

410 detail=ERROR_MESSAGES.FEATURE_DISABLED('Search query generation'), 

411 ) 

412 elif type == 'retrieval': 412 ↛ 413line 412 didn't jump to line 413 because the condition on line 412 was never true

413 if not await Config.get('task.query.retrieval.enable'): 

414 raise HTTPException( 

415 status_code=status.HTTP_400_BAD_REQUEST, 

416 detail=ERROR_MESSAGES.FEATURE_DISABLED('Query generation'), 

417 ) 

418 

419 if getattr(request.state, 'cached_queries', None): 419 ↛ 420line 419 didn't jump to line 420 because the condition on line 419 was never true

420 log.info('Reusing cached queries: %s', request.state.cached_queries) 

421 return request.state.cached_queries 

422 

423 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 423 ↛ 424line 423 didn't jump to line 424 because the condition on line 423 was never true

424 models = { 

425 **dict(request.app.state.MODELS.items()), 

426 request.state.model['id']: request.state.model, 

427 } 

428 else: 

429 models = request.app.state.MODELS 

430 

431 model_id = form_data['model'] 

432 if model_id not in models: 

433 raise HTTPException( 

434 status_code=status.HTTP_404_NOT_FOUND, 

435 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

436 ) 

437 

438 task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) 

439 

440 log.debug('generating %s queries using model %s for user %s', type, task_model_id, user.email) 

441 

442 query_template = await Config.get('task.query.prompt_template') 

443 if query_template.strip() != '': 

444 template = query_template 

445 else: 

446 template = DEFAULT_QUERY_GENERATION_PROMPT_TEMPLATE 

447 

448 content = await query_generation_template(template, form_data['messages'], user) 

449 

450 payload = { 

451 'model': task_model_id, 

452 'messages': [{'role': 'user', 'content': content}], 

453 'stream': False, 

454 'metadata': { 

455 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

456 'task': str(TASKS.QUERY_GENERATION), 

457 'task_body': form_data, 

458 'chat_id': form_data.get('chat_id', None), 

459 }, 

460 } 

461 

462 # Process the payload through the pipeline 

463 try: 

464 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

465 except Exception as e: 

466 raise e 

467 

468 payload = apply_task_model_params(payload, models, task_model_id, task_model_params) 

469 

470 try: 

471 return await generate_chat_completion(request, form_data=payload, user=user) 

472 except Exception as e: 

473 return JSONResponse( 

474 status_code=status.HTTP_400_BAD_REQUEST, 

475 content={'detail': str(e)}, 

476 ) 

477 

478 

479@router.post('/auto/completions') 

480async def generate_autocompletion(request: Request, form_data: dict, user=Depends(get_verified_user)): 

481 if not await Config.get('task.autocomplete.enable'): 

482 raise HTTPException( 

483 status_code=status.HTTP_400_BAD_REQUEST, 

484 detail=ERROR_MESSAGES.FEATURE_DISABLED('Autocompletion generation'), 

485 ) 

486 

487 type = form_data.get('type') 

488 prompt = form_data.get('prompt') 

489 messages = form_data.get('messages') 

490 

491 autocomplete_input_max_length = await Config.get('task.autocomplete.input_max_length') 

492 if autocomplete_input_max_length > 0: 492 ↛ 493line 492 didn't jump to line 493 because the condition on line 492 was never true

493 if len(prompt) > autocomplete_input_max_length: 

494 raise HTTPException( 

495 status_code=status.HTTP_400_BAD_REQUEST, 

496 detail=ERROR_MESSAGES.INPUT_TOO_LONG(autocomplete_input_max_length), 

497 ) 

498 

499 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 499 ↛ 500line 499 didn't jump to line 500 because the condition on line 499 was never true

500 models = { 

501 **dict(request.app.state.MODELS.items()), 

502 request.state.model['id']: request.state.model, 

503 } 

504 else: 

505 models = request.app.state.MODELS 

506 

507 model_id = form_data['model'] 

508 if model_id not in models: 

509 raise HTTPException( 

510 status_code=status.HTTP_404_NOT_FOUND, 

511 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

512 ) 

513 

514 task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) 

515 

516 log.debug('generating autocompletion using model %s for user %s', task_model_id, user.email) 

517 

518 autocomplete_template = await Config.get('task.autocomplete.prompt_template') 

519 if autocomplete_template.strip() != '': 

520 template = autocomplete_template 

521 else: 

522 template = DEFAULT_AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE 

523 

524 content = await autocomplete_generation_template(template, prompt, messages, type, user) 

525 

526 payload = { 

527 'model': task_model_id, 

528 'messages': [{'role': 'user', 'content': content}], 

529 'stream': False, 

530 'metadata': { 

531 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

532 'task': str(TASKS.AUTOCOMPLETE_GENERATION), 

533 'task_body': form_data, 

534 'chat_id': form_data.get('chat_id', None), 

535 }, 

536 } 

537 

538 # Process the payload through the pipeline 

539 try: 

540 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

541 except Exception as e: 

542 raise e 

543 

544 payload = apply_task_model_params(payload, models, task_model_id, task_model_params) 

545 

546 try: 

547 return await generate_chat_completion(request, form_data=payload, user=user) 

548 except Exception as e: 

549 log.error(f'Error generating chat completion: {e}') 

550 return JSONResponse( 

551 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, 

552 content={'detail': 'An internal error has occurred.'}, 

553 ) 

554 

555 

556@router.post('/emoji/completions') 

557async def generate_emoji(request: Request, form_data: dict, user=Depends(get_verified_user)): 

558 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 558 ↛ 559line 558 didn't jump to line 559 because the condition on line 558 was never true

559 models = { 

560 **dict(request.app.state.MODELS.items()), 

561 request.state.model['id']: request.state.model, 

562 } 

563 else: 

564 models = request.app.state.MODELS 

565 

566 model_id = form_data['model'] 

567 if model_id not in models: 

568 raise HTTPException( 

569 status_code=status.HTTP_404_NOT_FOUND, 

570 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

571 ) 

572 

573 task_model_id, _ = await get_task_model_generation_config(model_id, models) 

574 

575 log.debug('generating emoji using model %s for user %s ', task_model_id, user.email) 

576 

577 template = DEFAULT_EMOJI_GENERATION_PROMPT_TEMPLATE 

578 

579 content = await emoji_generation_template(template, form_data['prompt'], user) 

580 

581 payload = { 

582 'model': task_model_id, 

583 'messages': [{'role': 'user', 'content': content}], 

584 'stream': False, 

585 'metadata': { 

586 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

587 'task': str(TASKS.EMOJI_GENERATION), 

588 'task_body': form_data, 

589 'chat_id': form_data.get('chat_id', None), 

590 }, 

591 } 

592 

593 # Process the payload through the pipeline 

594 try: 

595 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

596 except Exception as e: 

597 raise e 

598 

599 payload = apply_task_model_params(payload, models, task_model_id, {'max_tokens': 4}) 

600 

601 try: 

602 return await generate_chat_completion(request, form_data=payload, user=user) 

603 except Exception as e: 

604 return JSONResponse( 

605 status_code=status.HTTP_400_BAD_REQUEST, 

606 content={'detail': str(e)}, 

607 ) 

608 

609 

610@router.post('/moa/completions') 

611async def generate_moa_response(request: Request, form_data: dict, user=Depends(get_verified_user)): 

612 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 612 ↛ 613line 612 didn't jump to line 613 because the condition on line 612 was never true

613 models = { 

614 **dict(request.app.state.MODELS.items()), 

615 request.state.model['id']: request.state.model, 

616 } 

617 else: 

618 models = request.app.state.MODELS 

619 

620 model_id = form_data['model'] 

621 

622 if model_id not in models: 

623 raise HTTPException( 

624 status_code=status.HTTP_404_NOT_FOUND, 

625 detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), 

626 ) 

627 

628 template = DEFAULT_MOA_GENERATION_PROMPT_TEMPLATE 

629 

630 content = moa_response_generation_template( 

631 template, 

632 form_data['prompt'], 

633 form_data['responses'], 

634 ) 

635 

636 payload = { 

637 'model': model_id, 

638 'messages': [{'role': 'user', 'content': content}], 

639 'stream': form_data.get('stream', False), 

640 'metadata': { 

641 **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 

642 'chat_id': form_data.get('chat_id', None), 

643 'task': str(TASKS.MOA_RESPONSE_GENERATION), 

644 'task_body': form_data, 

645 }, 

646 } 

647 

648 # Process the payload through the pipeline 

649 try: 

650 payload = await process_pipeline_inlet_filter(request, payload, user, models) 

651 except Exception as e: 

652 raise e 

653 

654 try: 

655 return await generate_chat_completion(request, form_data=payload, user=user) 

656 except Exception as e: 

657 return JSONResponse( 

658 status_code=status.HTTP_400_BAD_REQUEST, 

659 content={'detail': str(e)}, 

660 )