84 lines
3.0 KiB
Python
84 lines
3.0 KiB
Python
"""
|
||
JWT 工具:解析 Java 端签发的 token、从请求中提取 token。
|
||
与 Java 端 JwtService 共用同一个 HS256 密钥(AIIMAGE_JWT_SECRET)。
|
||
"""
|
||
import os
|
||
import sys
|
||
from typing import Optional, Tuple
|
||
|
||
from flask import request
|
||
|
||
try:
|
||
import jwt as pyjwt
|
||
_PYJWT_IMPORT_ERROR = None
|
||
except ImportError as _e:
|
||
pyjwt = None
|
||
_PYJWT_IMPORT_ERROR = _e
|
||
# 启动时立刻给出醒目提示,避免上线后才发现一直 401
|
||
print(
|
||
"[auth] WARNING: PyJWT 未安装,所有依赖 JWT 的接口都会返回 401。"
|
||
" 请执行 `pip install PyJWT==2.10.1` 或 `pip install -r requirements.txt`。",
|
||
file=sys.stderr,
|
||
)
|
||
|
||
|
||
_DEFAULT_SECRET = "please-change-this-secret-please-rotate-at-least-32-bytes"
|
||
COOKIE_NAME = os.getenv("AIIMAGE_AUTH_COOKIE_NAME", "aiimage_token")
|
||
|
||
|
||
def _signing_key() -> bytes:
|
||
"""与 Java JwtService.signingKey 保持一致:UTF-8 字节,不足 32 字节右侧补 0。"""
|
||
secret = os.getenv("AIIMAGE_JWT_SECRET", _DEFAULT_SECRET)
|
||
key_bytes = secret.encode("utf-8")
|
||
if len(key_bytes) < 32:
|
||
key_bytes = key_bytes + b"\x00" * (32 - len(key_bytes))
|
||
return key_bytes
|
||
|
||
|
||
def parse_token_with_reason(token: str) -> Tuple[Optional[dict], Optional[str]]:
|
||
"""解析 JWT,返回 (payload, error_reason)。
|
||
payload 命中时 error_reason 为 None;失败时 payload 为 None,error_reason 描述根因。"""
|
||
if not token:
|
||
return None, "token 为空"
|
||
if pyjwt is None:
|
||
return None, f"PyJWT 未安装({_PYJWT_IMPORT_ERROR})"
|
||
# 打印 token 头,便于发现 alg 不是 HS256 等情况(不验签,仅 base64 解码 header)
|
||
try:
|
||
header = pyjwt.get_unverified_header(token)
|
||
except Exception as e:
|
||
return None, f"无法解析 token header: {type(e).__name__}: {e}"
|
||
alg = header.get("alg") or "?"
|
||
try:
|
||
payload = pyjwt.decode(token, _signing_key(), algorithms=["HS256", "HS384", "HS512"])
|
||
except Exception as e:
|
||
# 常见:ExpiredSignatureError / InvalidSignatureError / DecodeError / InvalidAlgorithmError
|
||
return None, f"{type(e).__name__}: {e} (token alg={alg})"
|
||
sub = payload.get("sub")
|
||
try:
|
||
user_id = int(sub) if sub is not None else None
|
||
except (TypeError, ValueError):
|
||
return None, f"sub 字段非法: {sub!r}"
|
||
if user_id is None:
|
||
return None, "payload 缺少 sub 字段"
|
||
return {
|
||
"user_id": user_id,
|
||
"username": payload.get("username") or "",
|
||
"device_id": payload.get("deviceId") or "",
|
||
}, None
|
||
|
||
|
||
def parse_token(token: str) -> Optional[dict]:
|
||
"""解析 JWT,返回 {user_id, username, device_id};失败返回 None。"""
|
||
payload, _ = parse_token_with_reason(token)
|
||
return payload
|
||
|
||
|
||
def get_token_from_request() -> str:
|
||
"""优先从 Authorization: Bearer 取,其次从 cookie 取。"""
|
||
auth = request.headers.get("Authorization", "")
|
||
if auth.startswith("Bearer "):
|
||
token = auth[len("Bearer "):].strip()
|
||
if token:
|
||
return token
|
||
return (request.cookies.get(COOKIE_NAME) or "").strip()
|