chat.py 10.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187
  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 app.api import get_current_user, get_api_key
  8. from app.config.const import smart_message_error, http_400, http_500, http_200, complex_dialog_chat, \
  9. complex_knowledge_chat_deep, complex_knowledge_chat
  10. from app.models import UserModel
  11. from app.models.base_model import get_db
  12. from app.models.v2.chat import RetrievalRequest, ChatDataRequest, ComplexChatDao, SetModelRequest
  13. from app.models.v2.session_model import ChatData
  14. from app.service.v2.chat import service_chat_dialog, get_chat_info, service_chat_basic, \
  15. service_chat_workflow, service_chat_parameters, service_chat_sessions, service_chat_upload, \
  16. service_chat_sessions_list, service_chat_session_log, service_chunk_retrieval, service_complex_chat, \
  17. service_complex_upload, service_complex_model, service_get_complex_model
  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, chat_info.get_kb_ids()),
  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("/develop/{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("/retrieval")
  109. async def retrieve_chunks(request_data: RetrievalRequest, api_key: str = Depends(get_api_key)):
  110. 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)
  111. return {"records": records}
  112. @chat_router_v2.post("/complex/chat/completions")
  113. async def api_complex_chat_completions(chat: ChatDataRequest, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  114. # chat_mode = chat.chatMode
  115. if chat.isDeep == 2 and chat.chatMode == complex_knowledge_chat:
  116. chat.chatMode = complex_knowledge_chat_deep
  117. complex_chat = await ComplexChatDao(db).get_complex_chat_by_mode(chat.chatMode)
  118. if complex_chat:
  119. if not chat.sessionId:
  120. chat.sessionId = str(uuid.uuid4()).replace("-", "")
  121. return StreamingResponse(service_complex_chat(db, complex_chat.id, complex_chat.mode, current_user.id, chat),
  122. media_type="text/event-stream")
  123. else:
  124. error_msg = json.dumps(
  125. {"message": smart_message_error, "error": "\n**ERROR**: 网络异常,无法生成对话结果!", "status": http_500})
  126. return StreamingResponse(f"data: {error_msg}\n\n",
  127. media_type="text/event-stream")
  128. @chat_router_v2.post("/complex/upload/{chatMode}")
  129. async def api_complex_upload(chatMode:int, file: List[UploadFile] = File(...), current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  130. status_code = http_200
  131. complex_chat = await ComplexChatDao(db).get_complex_chat_by_mode(chatMode)
  132. if complex_chat:
  133. data = await service_complex_upload(db, complex_chat.id, file, current_user.id)
  134. if not data:
  135. status_code = http_400
  136. data = "{}"
  137. else:
  138. status_code = http_500
  139. data = "{}"
  140. return Response(data, media_type="application/json", status_code=status_code)
  141. @chat_router_v2.put("/complex/model")
  142. async def api_complex_model(chatModel:SetModelRequest, current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  143. status_code = http_200
  144. data = await service_complex_model(db, chatModel.chatType, chatModel.modelType, chatModel.modelName, chatModel.modelProvider)
  145. if data:
  146. status_code = http_500
  147. return Response(data, media_type="application/json", status_code=status_code)
  148. @chat_router_v2.get("/complex/model")
  149. async def api_get_complex_model(current_user: UserModel = Depends(get_current_user), db: Session = Depends(get_db)): # current_user: UserModel = Depends(get_current_user)
  150. status_code = http_200
  151. data = await service_get_complex_model(db)
  152. if not data:
  153. status_code = http_500
  154. return Response(data, media_type="application/json", status_code=status_code)