session.py 3.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293
  1. from sqlalchemy.orm import Session
  2. from Log import logger
  3. from app.models import AgentType
  4. from app.models.session_model import SessionModel
  5. class SessionService:
  6. def __init__(self, db: Session):
  7. self.db = db
  8. def create_session(self, session_id: str, name: str, agent_id: str, agent_type: AgentType, user_id: int) -> SessionModel:
  9. """
  10. 创建一个新的会话记录。
  11. 参数:
  12. session_id (str): 会话ID。
  13. name (str): 会话名称。
  14. agent_id (str): 代理ID。
  15. agent_type (AgentType): 代理类型。
  16. 返回:
  17. SessionModel: 新创建的会话模型实例,如果会话ID已存在则返回None。
  18. """
  19. logger.error("-------------xieru")
  20. existing_session = self.get_session_by_id(session_id)
  21. if existing_session:
  22. message=existing_session.message
  23. message.append({"role": "user", "content": name})
  24. self.update_session(session_id, message=message)
  25. logger.error("-------------xieru------------1")
  26. new_session = SessionModel(
  27. id=session_id,
  28. name=name[0:200],
  29. agent_id=agent_id,
  30. agent_type=agent_type,
  31. tenant_id = user_id,
  32. message=[{"role": "user", "content": name}]
  33. )
  34. self.db.add(new_session)
  35. self.db.commit()
  36. self.db.refresh(new_session)
  37. return new_session
  38. def get_session_by_id(self, session_id: str) -> SessionModel:
  39. """
  40. 根据会话ID获取会话记录。
  41. 参数:
  42. session_id (str): 会话ID。
  43. 返回:
  44. SessionModel: 查找到的会话模型实例,如果未找到则返回None。
  45. """
  46. return self.db.query(SessionModel).filter_by(id=session_id).first()
  47. def update_session(self, session_id: str, **kwargs) -> SessionModel:
  48. """
  49. 更新会话记录。
  50. 参数:
  51. session_id (str): 会话ID。
  52. kwargs: 需要更新的字段及其值。
  53. 返回:
  54. SessionModel: 更新后的会话模型实例。
  55. """
  56. logger.error("更新数据---------------------------")
  57. session = self.db.query(SessionModel).filter_by(id=session_id).first()
  58. if session:
  59. if "message" in kwargs:
  60. message = session.message
  61. logger.error(kwargs)
  62. message.append(kwargs["message"])
  63. session.message = message
  64. logger.error("更新数据--------------------------11111111-")
  65. logger.error(message)
  66. self.db.commit()
  67. self.db.refresh(session)
  68. logger.error("更新数据完成--------------------------1111111122222222-")
  69. return session
  70. def delete_session(self, session_id: str) -> None:
  71. """
  72. 删除会话记录。
  73. 参数:
  74. session_id (str): 会话ID。
  75. """
  76. session = self.get_session_by_id(session_id)
  77. if session:
  78. self.db.delete(session)
  79. self.db.commit()