token_model.py 1.8 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556
  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.models.base_model import Base
  6. class TokenModel(Base):
  7. __tablename__ = "token"
  8. id = Column(Integer, primary_key=True, index=True)
  9. user_id = Column(Integer, index=True)
  10. token = Column(Text(10000), unique=True, index=True)
  11. bisheng_token = Column(Text(10000), unique=True, index=True)
  12. ragflow_token = Column(Text(10000), unique=True, index=True)
  13. created_at = Column(DateTime, default=datetime.utcnow)
  14. def upsert_token(db: Session, user_id: int, access_token: str, bisheng_token: str, ragflow_token: str):
  15. # 参数验证
  16. if not isinstance(user_id, int) or user_id <= 0:
  17. return
  18. if not access_token or not bisheng_token or not ragflow_token:
  19. return
  20. db_token = None
  21. try:
  22. # 查询现有记录
  23. existing_token = db.query(TokenModel).filter_by(user_id=user_id).first()
  24. if existing_token:
  25. # 记录存在,进行更新
  26. existing_token.token = access_token
  27. existing_token.bisheng_token = bisheng_token
  28. existing_token.ragflow_token = ragflow_token
  29. else:
  30. # 记录不存在,进行插入
  31. db_token = TokenModel(
  32. user_id=user_id,
  33. token=access_token,
  34. bisheng_token=bisheng_token,
  35. ragflow_token=ragflow_token
  36. )
  37. db.add(db_token)
  38. # 提交事务
  39. db.commit()
  40. db.refresh(db_token)
  41. except Exception as e:
  42. # 异常处理
  43. db.rollback() # 回滚事务
  44. def get_token(db: Session, user_id: int) -> Type[TokenModel] | None:
  45. return db.query(TokenModel).filter_by(user_id=user_id).first()