user_context.py 4.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. """
  2. 请求级用户上下文 — Python contextvars 实现(等价于 Java ThreadLocal)
  3. 中间件在每个请求开始时解析 JWT(Cookie 或 Authorization Header),
  4. 将用户信息存入 ContextVar,后续任何位置通过 get_current_user() 即可获取,
  5. 无需显式传参或重复解码 JWT。
  6. 支持两种认证方式 + 自动续期:
  7. 1. Cookie: access_token(浏览器,HttpOnly 自动携带)
  8. 2. Header: Authorization: Bearer <token> + X-Refresh-Token(非浏览器设备)
  9. 3. 自动续期:非浏览器设备 access_token 过期时,中间件自动用 refresh_token 换新
  10. """
  11. from contextvars import ContextVar
  12. import jwt
  13. from starlette.middleware.base import BaseHTTPMiddleware
  14. from starlette.requests import Request
  15. from .jwt_utils import (
  16. verify_access_token, verify_refresh_token,
  17. create_access_token, create_refresh_token,
  18. )
  19. from .redis_service import (
  20. validate_refresh_token, revoke_refresh_token, store_refresh_token,
  21. )
  22. from .database import get_user_by_id
  23. # 核心:ContextVar 像 ThreadLocal,但兼容 asyncio
  24. # 每个请求只能看见自己的那份,请求之间完全隔离
  25. _current_user_var: ContextVar[dict | None] = ContextVar("current_user", default=None)
  26. def get_current_user() -> dict | None:
  27. """
  28. 获取当前登录用户。
  29. 返回值: {"id": int, "username": str} 或 None(未登录时)
  30. 无需 Request 参数,像全局变量一样调用。
  31. """
  32. return _current_user_var.get()
  33. def _extract_token(request: Request) -> str:
  34. """
  35. 从请求中提取 JWT Token,优先级:
  36. 1. Authorization: Bearer <token>(非浏览器设备)
  37. 2. Cookie: access_token(浏览器)
  38. """
  39. # 1. 检查 Authorization Header
  40. auth_header = request.headers.get("Authorization", "")
  41. if auth_header.startswith("Bearer "):
  42. return auth_header[7:].strip()
  43. # 2. 回退到 Cookie
  44. return request.cookies.get("access_token", "")
  45. def _try_auto_refresh(request: Request) -> tuple:
  46. """
  47. 当 access_token 过期时,尝试用 X-Refresh-Token 自动续期。
  48. 返回: (user_dict, new_access_token, new_refresh_token) 或 (None, None, None)
  49. """
  50. refresh_token = request.headers.get("X-Refresh-Token", "")
  51. if not refresh_token:
  52. return None, None, None
  53. try:
  54. payload = verify_refresh_token(refresh_token)
  55. stored = validate_refresh_token(payload["jti"])
  56. if stored is None or stored["user_id"] != payload["id"]:
  57. return None, None, None
  58. # 设备校验(非浏览器设备可能无 User-Agent,不阻塞)
  59. current_ua = request.headers.get("User-Agent", "")
  60. if stored.get("user_agent") and current_ua and stored["user_agent"] != current_ua:
  61. return None, None, None
  62. # 吊销旧 Token,签发新 Token
  63. revoke_refresh_token(payload["jti"])
  64. new_access = create_access_token(payload["id"])
  65. new_refresh, new_jti = create_refresh_token(payload["id"])
  66. store_refresh_token(payload["id"], new_jti, user_agent=current_ua)
  67. # 查用户信息
  68. user_info = get_user_by_id(payload["id"])
  69. if user_info:
  70. return {"id": user_info["id"], "username": user_info["username"]}, new_access, new_refresh
  71. except Exception:
  72. pass
  73. return None, None, None
  74. class UserContextMiddleware(BaseHTTPMiddleware):
  75. """
  76. FastAPI 中间件 — 自动解析 access_token(Cookie 或 Header),
  77. 注入当前用户到上下文。access_token 过期时自动续期(非浏览器设备)。
  78. """
  79. async def dispatch(self, request: Request, call_next):
  80. user = None
  81. new_access_token = None
  82. new_refresh_token = None
  83. token = _extract_token(request)
  84. if token:
  85. try:
  86. payload = verify_access_token(token)
  87. user_info = get_user_by_id(payload["id"])
  88. if user_info:
  89. user = {"id": user_info["id"], "username": user_info["username"]}
  90. except jwt.ExpiredSignatureError:
  91. # 过期 → 用 X-Refresh-Token 自动续期(对客户端透明)
  92. user, new_access_token, new_refresh_token = _try_auto_refresh(request)
  93. except Exception:
  94. pass # token 无效 → user 为 None,后续路由自行处理 401
  95. # 把当前用户 "set" 进上下文,类似 ThreadLocal.set()
  96. ctx_token = _current_user_var.set(user)
  97. try:
  98. response = await call_next(request)
  99. # 自动续期成功 → 通过响应头把新 Token 带回客户端
  100. if new_access_token:
  101. response.headers["X-Access-Token"] = new_access_token
  102. if new_refresh_token:
  103. response.headers["X-Refresh-Token"] = new_refresh_token
  104. return response
  105. finally:
  106. # 请求结束必须 reset,防止上下文泄漏到下一个请求
  107. _current_user_var.reset(ctx_token)