agent.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231
  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
  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. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  101. try:
  102. token = get_bisheng_token(db, current_user.id)
  103. result = await bisheng_service.get_session_log(token, agent_id, conversation_id)
  104. combined_logs = []
  105. last_question = None
  106. for session in result:
  107. print(session)
  108. # 检查 session 是否为 None
  109. if session is None:
  110. continue
  111. message = session.get('message')
  112. # 判断 message 是否是字符串,然后尝试解析为 JSON 对象
  113. if isinstance(message, str):
  114. try:
  115. message_json = json.loads(message)
  116. if 'question' in message_json:
  117. message = message_json['question']
  118. elif 'query' in message_json:
  119. message = message_json['query']
  120. elif 'report_name' in message_json:
  121. message = message_json['report_name']
  122. except json.JSONDecodeError:
  123. pass # 非 JSON 字符串,继续使用原始 message
  124. # 检查 message 是否为 None
  125. if message is None:
  126. continue
  127. if session.get('role') == 'question':
  128. last_question = message
  129. elif session.get('role') == 'answer' and last_question:
  130. combined_logs.append(last_question + "\n" + message)
  131. last_question = None
  132. return JSONResponse(status_code=200, content={"code": 200, "data": {"question":"", "answer": "\n".join(combined_logs)}})
  133. except Exception as e:
  134. raise HTTPException(status_code=500, detail=str(e))
  135. elif agent.agent_type == AgentType.BASIC:
  136. data = []
  137. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  138. if session:
  139. tmp_data = {}
  140. for i in session.log_to_json().get("message", []):
  141. if i.get("role") == "user":
  142. tmp_data["question"]=i.get("content")
  143. elif i.get("role") == "assistant":
  144. if isinstance(i.get("content"), dict):
  145. tmp_data["answer"] = i.get("content", {}).get("message")
  146. if "file_name" in i.get("content", {}):
  147. tmp_data["files"] = [{"file_name":i.get("content", {}).get("file_name"), "file_url":i.get("content", {}).get("file_url")}]
  148. else:
  149. tmp_data["answer"] = i.get("content")
  150. if "excel_url" in i:
  151. tmp_data["excel_url"] = i.get("excel_url")
  152. if "image_url" in i:
  153. tmp_data["image_url"] = i.get("image_url")
  154. if "sql" in i:
  155. tmp_data["sql"] = i.get("sql")
  156. if "code" in i:
  157. tmp_data["code"] = i.get("code")
  158. if "e" in i:
  159. tmp_data["e"] = i.get("e")
  160. if "image_name" in i:
  161. tmp_data["image_name"] = i.get("image_name")
  162. if "excel_name" in i:
  163. tmp_data["excel_name"] = i.get("excel_name")
  164. data.append(tmp_data)
  165. tmp_data = {}
  166. if tmp_data:
  167. data.append(tmp_data)
  168. return JSONResponse(status_code=200, content={"code": 200, "data": data})
  169. elif agent.agent_type == AgentType.DIFY:
  170. data = []
  171. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  172. if session:
  173. tmp_data = {}
  174. for i in session.log_to_json().get("message", []):
  175. if i.get("role") == "user":
  176. tmp_data["question"] = i.get("content")
  177. elif i.get("role") == "assistant":
  178. if isinstance(i.get("content"), dict):
  179. tmp_data["answer"] = i.get("content", {}).get("answer")
  180. if "file_name" in i.get("content", {}):
  181. tmp_data["files"] = [{"file_name": i.get("content", {}).get("file_name"),
  182. "file_url": i.get("content", {}).get("file_url")}]
  183. if "images" in i.get("content", {}):
  184. tmp_data["images"] = i.get("content", {}).get("images")
  185. else:
  186. tmp_data["answer"] = i.get("content")
  187. data.append(tmp_data)
  188. tmp_data = {}
  189. if tmp_data:
  190. data.append(tmp_data)
  191. return JSONResponse(status_code=200, content={"code": 200, "data": data})
  192. else:
  193. return JSONResponse(status_code=200, content={"code": 200, "log": "Unsupported agent type"})
  194. @router.get("/get-chat-id/{agent_id}", response_model=Response)
  195. async def get_chat_id(agent_id: str, db: Session = Depends(get_db)):
  196. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  197. if not agent:
  198. return Response(code=404, msg="Agent not found")
  199. return Response(code=200, msg="", data={"chat_id": uuid.uuid4().hex})