chat.py 7.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143
  1. import json
  2. import uuid
  3. from typing import List
  4. from fastapi import Depends, APIRouter, File, UploadFile
  5. from sqlalchemy.orm import Session
  6. from starlette.responses import StreamingResponse, Response
  7. from werkzeug.http import HTTP_STATUS_CODES
  8. from app.api import get_current_user, get_api_key
  9. from app.config.const import dialog_chat, advanced_chat, base_chat, agent_chat, workflow_chat, basic_chat, \
  10. smart_message_error, http_400, http_500, http_200
  11. from app.models import UserModel
  12. from app.models.base_model import get_db
  13. from app.models.v2.chat import RetrievalRequest
  14. from app.models.v2.session_model import ChatData
  15. from app.service.v2.chat import service_chat_dialog, get_chat_info, service_chat_basic, \
  16. service_chat_workflow, service_chat_parameters, service_chat_sessions, service_chat_upload, \
  17. service_chat_sessions_list, service_chat_session_log, service_chunk_retrieval, service_base_chunk_retrieval
  18. chat_router_v2 = APIRouter()
  19. # 对话
  20. @chat_router_v2.post("/chat/{chatId}/completions")
  21. async def api_chat_dialog(chatId:str, dialog: ChatData, current_user: UserModel = Depends(get_current_user),db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  22. chat_info = await get_chat_info(db, chatId)
  23. if not chat_info:
  24. error_msg = json.dumps(
  25. {"message": smart_message_error, "error": "\n**ERROR**: parameter exception", "status": http_400})
  26. return StreamingResponse(f"data: {error_msg}\n\n",
  27. media_type="text/event-stream")
  28. session_id = dialog.sessionId
  29. if not dialog.query:
  30. error_msg = json.dumps(
  31. {"message": smart_message_error, "error": "\n**ERROR**: question cannot be empty.", "status": http_400})
  32. return StreamingResponse(f"data: {error_msg}\n\n",
  33. media_type="text/event-stream")
  34. if not session_id:
  35. session = await service_chat_sessions(db, chatId, dialog.query)
  36. print(session)
  37. if not session or session.get("code") != 0:
  38. error_msg = json.dumps(
  39. {"message": smart_message_error, "error": "\n**ERROR**: chat agent error", "status": http_500})
  40. return StreamingResponse(f"data: {error_msg}\n\n",
  41. media_type="text/event-stream")
  42. session_id = session.get("data", {}).get("id")
  43. return StreamingResponse(service_chat_dialog(db, chatId, dialog.query, session_id, current_user.id, chat_info.mode),
  44. media_type="text/event-stream")
  45. @chat_router_v2.post("/agent/{chatId}/completions")
  46. async def api_chat_dialog(chatId:str, dialog: ChatData, current_user: UserModel = Depends(get_current_user),db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  47. chat_info = await get_chat_info(db, chatId)
  48. if not chat_info:
  49. error_msg = json.dumps(
  50. {"message": smart_message_error, "error": "\n**ERROR**: parameter exception", "status": http_400})
  51. return StreamingResponse(f"data: {error_msg}\n\n",
  52. media_type="text/event-stream")
  53. session_id = dialog.sessionId
  54. if not session_id:
  55. session_id = str(uuid.uuid4()).replace("-", "")
  56. return StreamingResponse(service_chat_workflow(db, chatId, dialog, session_id, current_user.id, chat_info.mode),
  57. media_type="text/event-stream")
  58. @chat_router_v2.post("/workflow/{chatId}/completions")
  59. async def api_chat_dialog(chatId:str, dialog: ChatData, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  60. chat_info = await get_chat_info(db, chatId)
  61. if not chat_info:
  62. error_msg = json.dumps(
  63. {"message": smart_message_error, "error": "\n**ERROR**: parameter exception", "status": http_400})
  64. return StreamingResponse(f"data: {error_msg}\n\n",
  65. media_type="text/event-stream")
  66. session_id = dialog.sessionId
  67. if not session_id:
  68. session_id = str(uuid.uuid4()).replace("-", "")
  69. return StreamingResponse(service_chat_workflow(db, chatId, dialog, session_id, current_user.id, chat_info.mode),
  70. media_type="text/event-stream")
  71. @chat_router_v2.post("/complex/{chatId}/completions")
  72. async def api_chat_dialog(chatId:str, dialog: ChatData, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  73. chat_info = await get_chat_info(db, chatId)
  74. if not chat_info:
  75. error_msg = json.dumps(
  76. {"message": smart_message_error, "error": "\n**ERROR**: parameter exception", "status": http_400})
  77. return StreamingResponse(f"data: {error_msg}\n\n",
  78. media_type="text/event-stream")
  79. session_id = dialog.sessionId
  80. if not session_id:
  81. session_id = str(uuid.uuid4()).replace("-", "")
  82. return StreamingResponse(service_chat_basic(db, chatId, dialog, session_id, current_user.id, chat_info.mode),
  83. media_type="text/event-stream")
  84. @chat_router_v2.get("/chat/{chatId}/parameters")
  85. async def api_chat_parameters(chatId:str, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  86. status_code = http_200
  87. data = await service_chat_parameters(db, chatId, current_user.id)
  88. if not data:
  89. status_code = http_400
  90. data = json.dumps({"code": http_400})
  91. return Response(data, media_type="application/json", status_code=status_code)
  92. @chat_router_v2.post("/{chatId}/upload")
  93. async def api_chat_upload(chatId:str, file: List[UploadFile] = File(...), current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  94. status_code = http_200
  95. data = await service_chat_upload(db, chatId, file, current_user.id)
  96. if not data:
  97. status_code = http_400
  98. data = "{}"
  99. return Response(data, media_type="application/json", status_code=status_code)
  100. @chat_router_v2.get("/chat/sessions")
  101. async def api_chat_sessions(chatId:str, current:int=1, pageSize:int=100, keyword:str="", current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  102. data = await service_chat_sessions_list(db, chatId, current, pageSize, current_user.id, keyword)
  103. return Response(data, media_type="application/json", status_code=http_200)
  104. @chat_router_v2.get("/chat/session_log")
  105. async def api_chat_sessions(sessionId:str, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  106. data = await service_chat_session_log(db, sessionId)
  107. return Response(data, media_type="application/json", status_code=http_200)
  108. # @chat_router_v2.post("/conversation/mindmap")
  109. # async def api_conversation_mindmap(chatId:str, current:int=1, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  110. # data = await service_chat_sessions_list(db, chatId, current, pageSize, current_user.id, keyword)
  111. # return Response(data, media_type="application/json", status_code=http_200)
  112. @chat_router_v2.post("/retrieval")
  113. async def retrieve_chunks(request_data: RetrievalRequest, api_key: str = Depends(get_api_key)):
  114. records = await service_chunk_retrieval(request_data.query, request_data.knowledge_id, request_data.retrieval_setting.top_k, request_data.retrieval_setting.score_threshold, api_key)
  115. return {"records": records}