token_model.py 2.8 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394
  1. from datetime import datetime
  2. from typing import Type
  3. from sqlalchemy import Column, Integer, DateTime, Text
  4. from sqlalchemy.orm import Session
  5. from app.config.const import RAGFLOW
  6. # from Log import logger
  7. from app.models.base_model import Base
  8. class TokenModel(Base):
  9. __tablename__ = "token"
  10. id = Column(Integer, primary_key=True, index=True)
  11. user_id = Column(Integer, index=True)
  12. token = Column(Text(10000))
  13. bisheng_token = Column(Text(10000))
  14. ragflow_token = Column(Text(10000))
  15. created_at = Column(DateTime, default=datetime.utcnow)
  16. def upsert_token(db: Session, user_id: int, access_token: str, bisheng_token: str, ragflow_token: str):
  17. # 参数验证
  18. if not isinstance(user_id, int) or user_id <= 0:
  19. return
  20. if not access_token or not bisheng_token or not ragflow_token:
  21. return
  22. db_token = None
  23. try:
  24. # 查询现有记录
  25. existing_token = db.query(TokenModel).filter_by(user_id=user_id).first()
  26. if existing_token:
  27. # 记录存在,进行更新
  28. existing_token.token = access_token
  29. existing_token.bisheng_token = bisheng_token
  30. existing_token.ragflow_token = ragflow_token
  31. else:
  32. # 记录不存在,进行插入
  33. db_token = TokenModel(
  34. user_id=user_id,
  35. token=access_token,
  36. bisheng_token=bisheng_token,
  37. ragflow_token=ragflow_token
  38. )
  39. db.add(db_token)
  40. # 提交事务
  41. db.commit()
  42. db.refresh(db_token)
  43. except Exception as e:
  44. # 异常处理
  45. db.rollback() # 回滚事务
  46. async def update_token(db: Session, user_id: int, access_token: str, token: dict):
  47. # 参数验证
  48. if not isinstance(user_id, int) or user_id <= 0:
  49. return
  50. db_token = None
  51. # print(token)
  52. try:
  53. # 查询现有记录
  54. db_token = db.query(TokenModel).filter_by(user_id=user_id).first()
  55. if db_token:
  56. # 记录存在,进行更新
  57. db_token.token = access_token
  58. db_token.ragflow_token = token.get(RAGFLOW, "")
  59. # db_token.token = access_token
  60. # for k, v in token.items():
  61. # setattr(db_token, k.replace("app", "token"), v)
  62. else:
  63. # 记录不存在,进行插入
  64. db_token = TokenModel(
  65. user_id=user_id,
  66. token=access_token,
  67. ragflow_token = token.get(RAGFLOW, "")
  68. )
  69. # for k, v in token.items():
  70. # setattr(db_token, k.replace("app", "token"), v)
  71. db.add(db_token)
  72. # 提交事务
  73. db.commit()
  74. db.refresh(db_token)
  75. except Exception as e:
  76. # logger.error(e)
  77. # 异常处理
  78. db.rollback() # 回滚事务