diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index c76529a..7de1c78 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -8,9 +8,45 @@ export const apiClient = axios.create({ timeout: 60000, }) -apiClient.interceptors.request.use((config) => { +// P1.4 token 静默刷新:JWT 剩余 <10min 时单飞调 /auth/refresh 换新, +// 长回测轮询不再因 60min 过期被 401 打断跳登录(401 拦截仍是兜底)。 +const REFRESH_AHEAD_SEC = 600 +let refreshPromise: Promise | null = null + +function tokenExpSec(token: string): number | null { + try { + const b64 = token.split('.')[1].replace(/-/g, '+').replace(/_/g, '/') + const payload = JSON.parse(atob(b64)) as { exp?: unknown } + return typeof payload.exp === 'number' ? payload.exp : null + } catch { + return null + } +} + +async function refreshToken(): Promise { + const auth = useAuthStore() + try { + // 裸 axios(不走 apiClient):避免拦截器递归/mock 短路 + const { data } = await axios.post('/api/v1/auth/refresh', null, { + headers: { Authorization: `Bearer ${auth.token}` }, + timeout: 10000, + }) + if (data?.token) auth.setToken(data.token as string, auth.username ?? '') + } catch { + // 刷新失败(如后端不可达):本次请求带旧 token 走,401 兜底处理 + } finally { + refreshPromise = null + } +} + +apiClient.interceptors.request.use(async (config) => { const auth = useAuthStore() if (auth.token) { + const exp = tokenExpSec(auth.token) + if (exp !== null && exp - Date.now() / 1000 < REFRESH_AHEAD_SEC) { + refreshPromise ??= refreshToken() + await refreshPromise + } config.headers.Authorization = `Bearer ${auth.token}` } diff --git a/sanguo_api/auth.py b/sanguo_api/auth.py index 2db600e..5689215 100644 --- a/sanguo_api/auth.py +++ b/sanguo_api/auth.py @@ -38,3 +38,18 @@ def verify_token(token: str) -> str: return payload["sub"] except jwt.PyJWTError: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无效 token") + + +def token_expires_in() -> int: + """当前配置的 token 有效期(秒),登录/刷新响应的 expires_in。""" + return int(_CONFIG["expire_minutes"] * 60) + + +def get_token_exp(token: str) -> int | None: + """解析 token 的 exp(unix 秒),不验签(供拦截器判断剩余时间);失败 → None。""" + try: + payload = jwt.decode(token, options={"verify_signature": False}) + exp = payload.get("exp") + return int(exp) if exp is not None else None + except (jwt.PyJWTError, ValueError, TypeError): + return None diff --git a/sanguo_api/routes.py b/sanguo_api/routes.py index a5bd2bb..e92e254 100644 --- a/sanguo_api/routes.py +++ b/sanguo_api/routes.py @@ -6,7 +6,11 @@ from fastapi import APIRouter, HTTPException, Depends, WebSocket, Query, Header from fastapi.responses import FileResponse from pydantic import BaseModel from .schemas import CtaBacktestRequest, OptimizeRequest, FactorAnalysisRequest -from .auth import verify_token as verify_token_impl, verify_password, create_token +from .auth import ( + create_token, token_expires_in, + verify_password, + verify_token as verify_token_impl, +) from .ws import manager from .strategy_registry import list_strategies, strategy_params, get_strategy_class from .kline import load_kline @@ -56,7 +60,19 @@ def login(req: LoginRequest): """Authenticate user and return JWT token""" if req.username != _auth_config["username"] or not verify_password(req.password, _auth_config["password_hash"]): raise HTTPException(status_code=401, detail="用户名或密码错误") - return {"token": create_token(req.username)} + return {"token": create_token(req.username), "expires_in": token_expires_in()} + + +@router.post("/auth/refresh", dependencies=[Depends(verify_token)]) +def refresh_token(authorization: str | None = Header(None)): + """P1.4 静默刷新:仍有效的旧 token 换新 token。 + + 过期 token 401(verify_token 拒)——不放过期续命;前端在剩余<10min 时 + 主动调本端点,长回测轮询不再因 60min 过期跳登录。单用户无 refresh + token 体系,滑动续期即够。 + """ + token = authorization.split(" ", 1)[1] + return {"token": create_token(verify_token_impl(token)), "expires_in": token_expires_in()} _BENCHMARKS = ("hs300", "zz500", "zz1000", "zz2000") diff --git a/tests/api/test_auth.py b/tests/api/test_auth.py index fbc2bb5..da965ef 100644 --- a/tests/api/test_auth.py +++ b/tests/api/test_auth.py @@ -22,3 +22,71 @@ def test_hash_and_verify_password(): assert h != "mypass" assert verify_password("mypass", h) is True assert verify_password("wrong", h) is False + + +# ---- P1.4 token 静默刷新 ---- + +def test_token_expiry_helpers(): + """get_token_exp 返回 unix 秒(供前端/拦截器判断剩余时间)。""" + import time + from sanguo_api.auth import create_token, get_token_exp, set_jwt_config + set_jwt_config(secret="test_secret", expire_minutes=60) + token = create_token("admin") + exp = get_token_exp(token) + assert exp is not None + assert abs(exp - (time.time() + 3600)) < 30 # ≈ now + 60min + + +def test_refresh_endpoint_returns_new_valid_token(): + """POST /auth/refresh:有效旧 token → 新 token(可过 verify),返回 expires_in。""" + from fastapi.testclient import TestClient + + from sanguo_api.app import create_app + from sanguo_api.auth import hash_password, set_jwt_config, verify_token + set_jwt_config(secret="test_secret", expire_minutes=60) + app = create_app(db_path=":memory:", auth_config={ + "username": "admin", + "password_hash": hash_password("pass123"), + "jwt_secret": "test_secret", + "expire_minutes": 60, + }) + client = TestClient(app) + old = client.post("/api/v1/auth/login", + json={"username": "admin", "password": "pass123"}).json()["token"] + r = client.post("/api/v1/auth/refresh", + headers={"Authorization": f"Bearer {old}"}) + assert r.status_code == 200 + data = r.json() + # 同秒签发的 JWT 可能逐字相同,以 exp 语义断言续期 + from sanguo_api.auth import get_token_exp + assert get_token_exp(data["token"]) >= get_token_exp(old) + assert verify_token(data["token"]) == "admin" + assert data["expires_in"] == 3600 + + +def test_refresh_endpoint_rejects_expired_or_missing_token(): + """过期/缺 token → 401(不放过期 token 无限续命)。""" + import time + + import jwt as pyjwt + from fastapi.testclient import TestClient + + from sanguo_api.app import create_app + from sanguo_api.auth import hash_password, set_jwt_config + set_jwt_config(secret="test_secret", expire_minutes=60) + app = create_app(db_path=":memory:", auth_config={ + "username": "admin", + "password_hash": hash_password("pass123"), + "jwt_secret": "test_secret", + "expire_minutes": 60, + }) + client = TestClient(app) + # 无 header + assert client.post("/api/v1/auth/refresh").status_code == 401 + # 过期 token(手工签一个已过期的) + expired = pyjwt.encode( + {"sub": "admin", "exp": int(time.time()) - 3600}, + "test_secret", algorithm="HS256") + r = client.post("/api/v1/auth/refresh", + headers={"Authorization": f"Bearer {expired}"}) + assert r.status_code == 401