redis_service.py 3.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105
  1. """Redis服务封装"""
  2. import json
  3. from datetime import timedelta
  4. import redis as redis_module
  5. import os
  6. REDIS_HOST = os.getenv("REDIS_HOST")
  7. REDIS_PORT = int(os.getenv("REDIS_PORT"))
  8. REDIS_PASSWORD = os.getenv("REDIS_PASSWORD")
  9. REDIS_DB = int(os.getenv("REDIS_DB"))
  10. # Refresh Token 过期时间(与JWT refresh token一致)
  11. REFRESH_TOKEN_TTL = timedelta(days=7)
  12. # Key 前缀
  13. PREFIX_REFRESH = "refresh_token:" # refresh_token:{jti} -> user_id
  14. PREFIX_USER_TOKENS = "user_tokens:" # user_tokens:{user_id} -> set of jti
  15. def get_redis() -> redis_module.Redis:
  16. """获取Redis连接"""
  17. return redis_module.Redis(
  18. host=REDIS_HOST,
  19. port=REDIS_PORT,
  20. password=REDIS_PASSWORD,
  21. db=REDIS_DB,
  22. decode_responses=True,
  23. socket_connect_timeout=3,
  24. )
  25. def store_refresh_token(user_id: int, jti: str, ttl_seconds: int = None, user_agent: str = "") -> bool:
  26. """
  27. 将 Refresh Token 的 jti 存入 Redis
  28. - refresh_token:{jti} -> json({"user_id": ..., "user_agent": ...}) (正向查:token -> user + 设备)
  29. - user_tokens:{user_id} -> set of jti (反向查:user -> tokens,用于踢下线)
  30. """
  31. if ttl_seconds is None:
  32. ttl_seconds = int(REFRESH_TOKEN_TTL.total_seconds())
  33. r = get_redis()
  34. try:
  35. pipe = r.pipeline()
  36. data = json.dumps({"user_id": user_id, "user_agent": user_agent})
  37. pipe.setex(f"{PREFIX_REFRESH}{jti}", ttl_seconds, data)
  38. pipe.sadd(f"{PREFIX_USER_TOKENS}{user_id}", jti)
  39. pipe.expire(f"{PREFIX_USER_TOKENS}{user_id}", ttl_seconds)
  40. pipe.execute()
  41. return True
  42. finally:
  43. r.close()
  44. def validate_refresh_token(jti: str) -> dict | None:
  45. """
  46. 验证 Refresh Token 是否在 Redis 中有效
  47. 返回值: {"user_id": int, "user_agent": str} 或 None
  48. """
  49. r = get_redis()
  50. try:
  51. data = r.get(f"{PREFIX_REFRESH}{jti}")
  52. if data is None:
  53. return None
  54. parsed = json.loads(data)
  55. return {
  56. "user_id": int(parsed["user_id"]),
  57. "user_agent": parsed.get("user_agent", ""),
  58. }
  59. finally:
  60. r.close()
  61. def revoke_refresh_token(jti: str) -> bool:
  62. """吊销单个 Refresh Token"""
  63. r = get_redis()
  64. try:
  65. data = r.get(f"{PREFIX_REFRESH}{jti}")
  66. if data is None:
  67. return False
  68. parsed = json.loads(data)
  69. uid = int(parsed["user_id"])
  70. pipe = r.pipeline()
  71. pipe.delete(f"{PREFIX_REFRESH}{jti}")
  72. pipe.srem(f"{PREFIX_USER_TOKENS}{uid}", jti)
  73. pipe.execute()
  74. return True
  75. finally:
  76. r.close()
  77. def revoke_all_user_tokens(user_id: int) -> int:
  78. """吊销用户的所有 Refresh Token(全部踢下线)"""
  79. r = get_redis()
  80. try:
  81. jtis = r.smembers(f"{PREFIX_USER_TOKENS}{user_id}")
  82. if not jtis:
  83. return 0
  84. pipe = r.pipeline()
  85. for jti in jtis:
  86. pipe.delete(f"{PREFIX_REFRESH}{jti}")
  87. pipe.delete(f"{PREFIX_USER_TOKENS}{user_id}")
  88. pipe.execute()
  89. return len(jtis)
  90. finally:
  91. r.close()