session_model.py 5.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155
  1. import json
  2. import pytz
  3. from datetime import datetime
  4. from sqlalchemy.orm import Session
  5. from typing import Optional, Type
  6. from pydantic import BaseModel
  7. from sqlalchemy import Column, String, Integer, DateTime, JSON, TEXT, Index
  8. from Log import logger
  9. from app.models.agent_model import AgentType
  10. from app.models.base_model import Base
  11. def current_time():
  12. tz = pytz.timezone('Asia/Shanghai')
  13. return datetime.now(tz)
  14. class ChatSessionModel(Base):
  15. __tablename__ = "chat_sessions"
  16. # __table_args__ = (
  17. # Index('idx_username', 'username'),
  18. # )
  19. id = Column(Integer, primary_key=True)
  20. name = Column(String(255))
  21. agent_id = Column(String(255))
  22. agent_type = Column(Integer) # 目前只存basic的,ragflow和bisheng的调接口获取
  23. create_date = Column(DateTime, default=current_time) # 创建时间,默认值为当前时区时间
  24. update_date = Column(DateTime, default=current_time, onupdate=current_time, index=True) # 更新时间,默认值为当前时区时间,更新时自动更新
  25. tenant_id = Column(Integer) # 创建人
  26. message = Column(TEXT) # 说明
  27. reference = Column(TEXT) # 说明
  28. conversation_id = Column(String(64))
  29. session_id = Column(String(36), index=True)
  30. chat_mode = Column(Integer)
  31. # to_dict 方法
  32. def to_dict(self):
  33. return {
  34. 'id': self.id,
  35. 'name': self.name,
  36. 'agent_type': self.agent_type,
  37. 'agent_id': self.agent_id,
  38. 'create_date': self.create_date.strftime("%Y-%m-%d %H:%M:%S"),
  39. 'update_date': self.update_date.strftime("%Y-%m-%d %H:%M:%S"),
  40. }
  41. def log_to_json(self):
  42. return {
  43. 'id': self.id,
  44. 'name': self.name,
  45. 'agent_type': self.agent_type,
  46. 'agent_id': self.agent_id,
  47. 'create_date': self.create_date.strftime("%Y-%m-%d %H:%M:%S"),
  48. 'update_date': self.update_date.strftime("%Y-%m-%d %H:%M:%S"),
  49. 'message': json.loads(self.message)
  50. }
  51. def add_message(self, message: dict):
  52. if self.message is None:
  53. self.message = '[]'
  54. try:
  55. msg = json.loads(self.message)
  56. msg.append(message)
  57. except Exception as e:
  58. return
  59. self.message = json.dumps(msg)
  60. class ChatDialogData(BaseModel):
  61. sessionId: Optional[str] = ""
  62. question: str
  63. chatId: str
  64. class ChatSessionDao:
  65. def __init__(self, db: Session):
  66. self.db = db
  67. def create_session(self, session_id: str, name: str, agent_id: str, agent_type: int, user_id: int, message: str,reference:str) -> ChatSessionModel:
  68. new_session = ChatSessionModel(
  69. id=session_id,
  70. name=name[0:255],
  71. agent_id=agent_id,
  72. agent_type=agent_type,
  73. create_date=current_time(),
  74. update_date=current_time(),
  75. tenant_id=user_id,
  76. message=message,
  77. reference=reference,
  78. )
  79. self.db.add(new_session)
  80. self.db.commit()
  81. self.db.refresh(new_session)
  82. return new_session
  83. def get_session_by_id(self, session_id: str) -> Type[ChatSessionModel] | None:
  84. session = self.db.query(ChatSessionModel).filter_by(id=session_id).first()
  85. if session and session.message is None:
  86. session.message = '[]'
  87. return session
  88. def update_session_by_id(self, session_id: str, **kwargs) -> Type[ChatSessionModel] | None:
  89. session = self.get_session_by_id(session_id)
  90. if session:
  91. if "message" in kwargs:
  92. session.add_message(kwargs["message"])
  93. # 替换其他字段
  94. for key, value in kwargs.items():
  95. if key != "message":
  96. setattr(session, key, value)
  97. session.update_date = current_time()
  98. try:
  99. self.db.commit()
  100. self.db.refresh(session)
  101. except Exception as e:
  102. logger.error(e)
  103. self.db.rollback()
  104. return session
  105. def create_session(self, session_id: str, name: str, agent_id: str, agent_type: AgentType, user_id: int) -> ChatSessionModel:
  106. existing_session = self.get_session_by_id(session_id)
  107. if existing_session:
  108. existing_session.add_message({"role": "user", "content": name})
  109. existing_session.update_date = current_time()
  110. self.db.commit()
  111. self.db.refresh(existing_session)
  112. return existing_session
  113. new_session = ChatSessionModel(
  114. id=session_id,
  115. name=name[0:50],
  116. agent_id=agent_id,
  117. agent_type=agent_type,
  118. tenant_id=user_id,
  119. message=json.dumps([{"role": "user", "content": name}])
  120. )
  121. self.db.add(new_session)
  122. self.db.commit()
  123. self.db.refresh(new_session)
  124. return new_session
  125. def delete_session(self, session_id: str) -> None:
  126. """
  127. 删除会话记录。
  128. 参数:
  129. session_id (str): 会话ID。
  130. """
  131. session = self.get_session_by_id(session_id)
  132. if session:
  133. self.db.delete(session)
  134. self.db.commit()