agent.py 7.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155
  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.config.config import settings
  9. from app.models.agent_model import AgentType, AgentModel
  10. from app.models.base_model import get_db
  11. from app.models.session_model import SessionModel
  12. from app.models.user_model import UserModel
  13. from app.service.bisheng import BishengService
  14. from app.service.dialog import get_session_history
  15. from app.service.ragflow import RagflowService
  16. from app.service.service_token import get_ragflow_token, get_bisheng_token
  17. router = APIRouter()
  18. @router.get("/list", response_model=ResponseList)
  19. async def agent_list(db: Session = Depends(get_db)):
  20. agents = db.query(AgentModel).order_by(AgentModel.sort.asc()).all()
  21. result = [item.to_dict() for item in agents]
  22. return ResponseList(code=200, msg="", data=result)
  23. @router.get("/{agent_id}/sessions", response_model=ResponseList)
  24. async def chat_list(
  25. agent_id: str,
  26. page: int = Query(1, ge=1),
  27. limit: int = Query(1000, ge=1, le=1000),
  28. db: Session = Depends(get_db),
  29. current_user: UserModel = Depends(get_current_user)):
  30. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  31. if not agent:
  32. return ResponseList(code=404, msg="Agent not found")
  33. if agent.agent_type == AgentType.RAGFLOW:
  34. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  35. try:
  36. token = get_ragflow_token(db, current_user.id)
  37. result = await ragflow_service.get_chat_sessions(token, agent_id)
  38. if not result:
  39. result = await get_session_history(db, current_user.id, agent_id)
  40. except Exception as e:
  41. raise HTTPException(status_code=500, detail=str(e))
  42. return ResponseList(code=200, msg="", data=result)
  43. elif agent.agent_type == AgentType.BISHENG:
  44. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  45. try:
  46. token = get_bisheng_token(db, current_user.id)
  47. result = await bisheng_service.get_chat_sessions(token, page, limit)
  48. except Exception as e:
  49. raise HTTPException(status_code=500, detail=str(e))
  50. return ResponseList(code=200, msg="", data=result)
  51. elif agent.agent_type == AgentType.BASIC:
  52. offset = (page - 1) * limit
  53. 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()
  54. result = [item.to_dict() for item in records]
  55. return ResponseList(code=200, msg="", data=result)
  56. else:
  57. return ResponseList(code=200, msg="Unsupported agent type")
  58. @router.get("/{agent_id}/{conversation_id}/session_log")
  59. async def session_log(agent_id: str, conversation_id: str, db: Session = Depends(get_db), current_user: UserModel = Depends(get_current_user)):
  60. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  61. if not agent:
  62. return Response(code=404, msg="Agent not found")
  63. if agent.agent_type == AgentType.RAGFLOW:
  64. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  65. try:
  66. token = get_ragflow_token(db, current_user.id)
  67. result = await ragflow_service.get_session_log(token, conversation_id)
  68. if 'session_log' in result and 'reference' in result:
  69. combined_logs = []
  70. last_question = None
  71. references = result['reference']
  72. reference_index = 0
  73. for session in result['session_log']:
  74. if session['role'] == 'user':
  75. last_question = session['message']
  76. elif session['role'] == 'assistant' and last_question:
  77. if reference_index < len(references):
  78. reference = references[reference_index]
  79. else:
  80. reference = None
  81. combined_logs.append({
  82. 'question': last_question,
  83. 'answer': session['message'],
  84. 'reference': reference
  85. })
  86. last_question = None
  87. reference_index += 1
  88. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs})
  89. else:
  90. return JSONResponse(status_code=200, content={"code": 400, "message": "Invalid result structure"})
  91. except Exception as e:
  92. raise HTTPException(status_code=500, detail=str(e))
  93. elif agent.agent_type == AgentType.BISHENG:
  94. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  95. try:
  96. token = get_bisheng_token(db, current_user.id)
  97. result = await bisheng_service.get_session_log(token, agent_id, conversation_id)
  98. combined_logs = []
  99. last_question = None
  100. for session in result:
  101. message = session['message']
  102. # 判断message是字符串还是json 对象,如果是json取其中的question字段,或者report_name字段赋值给message
  103. try:
  104. message_json = json.loads(message)
  105. if 'question' in message_json:
  106. message = message_json['question']
  107. elif 'query' in message_json:
  108. message = message_json['query']
  109. elif 'report_name' in message_json:
  110. message = message_json['report_name']
  111. except json.JSONDecodeError:
  112. pass
  113. if session['role'] == 'question':
  114. last_question = message
  115. elif session['role'] == 'answer' and last_question:
  116. combined_logs.append({
  117. 'question': last_question,
  118. 'answer': message
  119. })
  120. last_question = None
  121. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs})
  122. except Exception as e:
  123. raise HTTPException(status_code=500, detail=str(e))
  124. elif agent.agent_type == AgentType.BASIC:
  125. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  126. return JSONResponse(status_code=200, content={"code": 200, "data": session.log_to_json() if session else {}})
  127. else:
  128. return JSONResponse(status_code=200, content={"code": 200, "log": "Unsupported agent type"})
  129. @router.get("/get-chat-id/{agent_id}", response_model=Response)
  130. async def get_chat_id(agent_id: str, db: Session = Depends(get_db)):
  131. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  132. if not agent:
  133. return Response(code=404, msg="Agent not found")
  134. return Response(code=200, msg="", data={"chat_id": uuid.uuid4().hex})