token_model.py 2.9 KB

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