session.py 3.3 KB

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