agent.py 12 KB

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