agent.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260
  1. import json
  2. import uuid
  3. from fastapi import Depends, APIRouter, Query, HTTPException
  4. from fastapi.responses import JSONResponse
  5. from pydantic import BaseModel
  6. from sqlalchemy.orm import Session
  7. from app.api import Response, get_current_user, ResponseList, process_files
  8. from app.api.user import reset_user_pwd
  9. from app.config.config import settings
  10. from app.models.agent_model import AgentType, AgentModel
  11. from app.models.base_model import get_db
  12. from app.models.session_model import SessionModel
  13. from app.models.user_model import UserModel
  14. from app.service.bisheng import BishengService
  15. from app.service.dialog import get_session_history
  16. from app.service.ragflow import RagflowService
  17. from app.service.service_token import get_ragflow_token, get_bisheng_token
  18. router = APIRouter()
  19. @router.get("/list", response_model=ResponseList)
  20. async def agent_list(db: Session = Depends(get_db)):
  21. agents = db.query(AgentModel).order_by(AgentModel.sort.asc()).all()
  22. result = [item.to_dict() for item in agents]
  23. return ResponseList(code=200, msg="", data=result)
  24. @router.get("/{agent_id}/sessions", response_model=ResponseList)
  25. async def chat_list(
  26. agent_id: str,
  27. page: int = Query(1, ge=1),
  28. limit: int = Query(1000, ge=1, le=1000),
  29. db: Session = Depends(get_db),
  30. current_user: UserModel = Depends(get_current_user)):
  31. print(111)
  32. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  33. if not agent:
  34. return ResponseList(code=404, msg="Agent not found")
  35. if agent.agent_type == AgentType.RAGFLOW:
  36. print(222)
  37. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  38. try:
  39. token = await get_ragflow_token(db, current_user.id)
  40. result = await ragflow_service.get_chat_sessions(token, agent_id)
  41. if not result:
  42. result = await get_session_history(db, current_user.id, agent_id)
  43. except Exception as e:
  44. print(e)
  45. raise HTTPException(status_code=500, detail=str(e))
  46. return ResponseList(code=200, msg="", data=result)
  47. elif agent.agent_type == AgentType.BISHENG:
  48. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  49. try:
  50. token = await get_bisheng_token(db, current_user.id)
  51. result = await bisheng_service.get_chat_sessions(token, agent_id, page, limit)
  52. except Exception as e:
  53. raise HTTPException(status_code=500, detail=str(e))
  54. return ResponseList(code=200, msg="", data=result)
  55. elif agent.agent_type == AgentType.BASIC:
  56. offset = (page - 1) * limit
  57. records = db.query(SessionModel).filter(SessionModel.agent_id == agent_id, SessionModel.tenant_id==current_user.id).order_by(SessionModel.create_date.desc()).offset(offset).limit(limit).all()
  58. result = [item.to_dict() for item in records]
  59. return ResponseList(code=200, msg="", data=result)
  60. elif agent.agent_type == AgentType.DIFY:
  61. offset = (page - 1) * limit
  62. records = db.query(SessionModel).filter(SessionModel.agent_id == agent_id, SessionModel.tenant_id==current_user.id).order_by(SessionModel.create_date.desc()).offset(offset).limit(limit).all()
  63. result = [item.to_dict() for item in records]
  64. return ResponseList(code=200, msg="", data=result)
  65. else:
  66. return ResponseList(code=200, msg="Unsupported agent type")
  67. @router.get("/{agent_id}/{conversation_id}/session_log")
  68. async def session_log(agent_id: str, conversation_id: str, db: Session = Depends(get_db), current_user: UserModel = Depends(get_current_user)):
  69. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  70. if not agent:
  71. return Response(code=404, msg="Agent not found")
  72. if agent.agent_type == AgentType.RAGFLOW:
  73. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  74. try:
  75. token = await get_ragflow_token(db, current_user.id)
  76. result = await ragflow_service.get_session_log(token, conversation_id)
  77. if 'session_log' in result and 'reference' in result:
  78. combined_logs = []
  79. last_question = None
  80. references = result['reference']
  81. reference_index = 0
  82. for session in result['session_log']:
  83. if session['role'] == 'user':
  84. last_question = session['message']
  85. elif session['role'] == 'assistant' and last_question:
  86. if reference_index < len(references):
  87. reference = references[reference_index]
  88. else:
  89. reference = None
  90. combined_logs.append({
  91. 'question': last_question,
  92. 'answer': session['message'],
  93. 'reference': reference
  94. })
  95. last_question = None
  96. reference_index += 1
  97. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs})
  98. else:
  99. return JSONResponse(status_code=200, content={"code": 400, "message": "Invalid result structure"})
  100. except Exception as e:
  101. raise HTTPException(status_code=500, detail=str(e))
  102. elif agent.agent_type == AgentType.BISHENG:
  103. is_join = False
  104. if agent.name == "报告生成":
  105. is_join = True
  106. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  107. try:
  108. token = await get_bisheng_token(db, current_user.id)
  109. result = await bisheng_service.get_session_log(token, agent_id, conversation_id)
  110. combined_logs = []
  111. last_question = None
  112. answer_str = ""
  113. files = []
  114. for session in result:
  115. # print(session)
  116. # 检查 session 是否为 None
  117. if session is None:
  118. continue
  119. message = session.get('message')
  120. # 判断 message 是否是字符串,然后尝试解析为 JSON 对象
  121. if isinstance(message, str):
  122. try:
  123. message_json = json.loads(message)
  124. if 'question' in message_json:
  125. message = message_json['question']
  126. elif 'query' in message_json:
  127. message = message_json['query']
  128. elif 'report_name' in message_json:
  129. message = message_json['report_name']
  130. except Exception as e:
  131. pass # 非 JSON 字符串,继续使用原始 message
  132. if session.get('files') and isinstance(session.get('files'), str):
  133. try:
  134. files = json.loads(session.get('files'))
  135. process_files(files, agent_id)
  136. except Exception as e:
  137. pass # 非 JSON 字符串,继续使用原始 message
  138. # 检查 message 是否为 None
  139. if message is None:
  140. continue
  141. if is_join:
  142. ...
  143. if session.get('role') == 'question':
  144. last_question = message
  145. elif session.get('role') == 'answer':
  146. answer_str += message
  147. else:
  148. if session.get('role') == 'question':
  149. last_question = message
  150. elif session.get('role') == 'answer' and last_question:
  151. combined_logs.append({
  152. 'question': last_question,
  153. 'answer': message
  154. })
  155. last_question = None
  156. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs if combined_logs else [{'question': last_question,
  157. 'answer': answer_str, 'files': files}]})
  158. except Exception as e:
  159. raise HTTPException(status_code=500, detail=str(e))
  160. elif agent.agent_type == AgentType.BASIC:
  161. data = []
  162. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  163. if session:
  164. tmp_data = {}
  165. for i in session.log_to_json().get("message", []):
  166. if i.get("role") == "user":
  167. tmp_data["question"]=i.get("content")
  168. elif i.get("role") == "assistant":
  169. if isinstance(i.get("content"), dict):
  170. tmp_data["answer"] = i.get("content", {}).get("message")
  171. if "file_name" in i.get("content", {}):
  172. tmp_data["files"] = [{"file_name":i.get("content", {}).get("file_name"), "file_url":i.get("content", {}).get("file_url")}]
  173. else:
  174. tmp_data["answer"] = i.get("content")
  175. if "excel_url" in i:
  176. tmp_data["excel_url"] = i.get("excel_url")
  177. if "image_url" in i:
  178. tmp_data["image_url"] = i.get("image_url")
  179. if "sql" in i:
  180. tmp_data["sql"] = i.get("sql")
  181. if "code" in i:
  182. tmp_data["code"] = i.get("code")
  183. if "e" in i:
  184. tmp_data["e"] = i.get("e")
  185. if "image_name" in i:
  186. tmp_data["image_name"] = i.get("image_name")
  187. if "excel_name" in i:
  188. tmp_data["excel_name"] = i.get("excel_name")
  189. data.append(tmp_data)
  190. tmp_data = {}
  191. if tmp_data:
  192. data.append(tmp_data)
  193. return JSONResponse(status_code=200, content={"code": 200, "data": data})
  194. elif agent.agent_type == AgentType.DIFY:
  195. data = []
  196. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  197. if session:
  198. tmp_data = {}
  199. for i in session.log_to_json().get("message", []):
  200. if i.get("role") == "user":
  201. tmp_data["question"] = i.get("content")
  202. elif i.get("role") == "assistant":
  203. if isinstance(i.get("content"), dict):
  204. tmp_data["answer"] = i.get("content", {}).get("answer")
  205. if "file_name" in i.get("content", {}):
  206. tmp_data["files"] = [{"file_name": i.get("content", {}).get("file_name"),
  207. "file_url": i.get("content", {}).get("file_url")}]
  208. if "images" in i.get("content", {}):
  209. tmp_data["images"] = i.get("content", {}).get("images")
  210. if "download_url" in i.get("content", {}):
  211. tmp_data["download_url"] = i.get("content", {}).get("download_url")
  212. else:
  213. tmp_data["answer"] = i.get("content")
  214. data.append(tmp_data)
  215. tmp_data = {}
  216. if tmp_data:
  217. data.append(tmp_data)
  218. return JSONResponse(status_code=200, content={"code": 200, "data": data})
  219. else:
  220. return JSONResponse(status_code=200, content={"code": 200, "log": "Unsupported agent type"})
  221. @router.get("/get-chat-id/{agent_id}", response_model=Response)
  222. async def get_chat_id(agent_id: str, db: Session = Depends(get_db)):
  223. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  224. if not agent:
  225. return Response(code=404, msg="Agent not found")
  226. return Response(code=200, msg="", data={"chat_id": uuid.uuid4().hex})