session.py 2.7 KB

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