session.py 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112
  1. import json
  2. from typing import Type
  3. from sqlalchemy.orm import Session
  4. from Log import logger
  5. from app.models import AgentType, current_time
  6. from app.models.session_model import SessionModel
  7. class SessionService:
  8. def __init__(self, db: Session):
  9. self.db = db
  10. def create_session(self, session_id: str, name: str, agent_id: str, agent_type: AgentType, user_id: int, message: dict = None, workflow_type: int = 0) -> Type[
  11. SessionModel] | SessionModel:
  12. """
  13. 创建一个新的会话记录。
  14. 参数:
  15. session_id (str): 会话ID。
  16. name (str): 会话名称。
  17. agent_id (str): 代理ID。
  18. agent_type (AgentType): 代理类型。
  19. 返回:
  20. SessionModel: 新创建的会话模型实例,如果会话ID已存在则返回None。
  21. """
  22. if not message:
  23. message = {"role": "user", "content": name}
  24. existing_session = self.get_session_by_id(session_id)
  25. if existing_session:
  26. print("update success")
  27. # existing_session.add_message({"role": "user", "content": name})
  28. existing_session.add_message(message)
  29. existing_session.update_date = current_time()
  30. self.db.commit()
  31. self.db.refresh(existing_session)
  32. return existing_session
  33. new_session = SessionModel(
  34. id=session_id,
  35. name=name[0:50],
  36. agent_id=agent_id,
  37. agent_type=agent_type,
  38. tenant_id=user_id,
  39. workflow=workflow_type,
  40. message = json.dumps([message])
  41. )
  42. self.db.add(new_session)
  43. self.db.commit()
  44. self.db.refresh(new_session)
  45. return new_session
  46. def get_session_by_id(self, session_id: str) -> Type[SessionModel] | None:
  47. """
  48. 根据会话ID获取会话记录。
  49. 参数:
  50. session_id (str): 会话ID。
  51. 返回:
  52. SessionModel: 查找到的会话模型实例,如果未找到则返回None。
  53. """
  54. session = self.db.query(SessionModel).filter_by(id=session_id).first()
  55. if session and session.message is None:
  56. session.message = '[]'
  57. return session
  58. def update_session(self, session_id: str, **kwargs) -> Type[SessionModel] | None:
  59. """
  60. 更新会话记录。
  61. 参数:
  62. session_id (str): 会话ID。
  63. kwargs: 需要更新的字段及其值。
  64. 返回:
  65. SessionModel: 更新后的会话模型实例。
  66. """
  67. logger.error("更新数据---------------------------")
  68. self.db.commit()
  69. session = self.get_session_by_id(session_id)
  70. if session:
  71. if "message" in kwargs:
  72. session.add_message(kwargs["message"])
  73. # 替换其他字段
  74. for key, value in kwargs.items():
  75. if key != "message":
  76. setattr(session, key, value)
  77. session.update_date = current_time()
  78. try:
  79. self.db.commit()
  80. self.db.refresh(session)
  81. except Exception as e:
  82. logger.error(e)
  83. self.db.rollback()
  84. return session
  85. def delete_session(self, session_id: str) -> None:
  86. """
  87. 删除会话记录。
  88. 参数:
  89. session_id (str): 会话ID。
  90. """
  91. session = self.get_session_by_id(session_id)
  92. if session:
  93. self.db.delete(session)
  94. self.db.commit()