Files
backend/backend/app/core/token_store.py
T
34047007@qq.com a6cd99a4ca
CI / backend (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
feat: initial commit - oncology literature search platform
OncoLit: a multi-tenant oncology literature search, feed, and
collaboration platform. Built with FastAPI + Vue 3 + PostgreSQL.
Includes PubMed pipeline, drug approvals, AI summaries, and
systematic review tools.
2026-07-27 07:59:18 +08:00

198 lines
7.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""JWT Token 持久存储 — Redis 黑名单 + Refresh Token 管理
三层 fallbackRedis → JSON 文件(dev 模式) → 内存
文件持久化确保 --reload 等重启场景下 refresh token 不丢失。
"""
import json
import os
import time
from pathlib import Path
from app.config import settings
class TokenStore:
"""Refresh Token 存储 + 黑名单(Redis → 文件 → 内存三层 fallback"""
_FILE_PATH = Path(__file__).resolve().parent.parent.parent / "data" / "token_store.json"
def __init__(self):
self._redis = None
self._fallback_blacklist: dict[str, float] = {}
self._fallback_refresh: dict[str, dict] = {}
# 文件最后加载时间戳,避免重复读盘
self._last_file_load = 0.0
self._file_dirty = False
# ── 文件持久化 ──
def _ensure_file_dir(self):
self._FILE_PATH.parent.mkdir(parents=True, exist_ok=True)
def _load_file(self):
"""从 JSON 文件加载 token 数据到内存(不覆盖已有内存数据)。"""
now = time.time()
if now - self._last_file_load < 1.0:
return
self._last_file_load = now
try:
if self._FILE_PATH.exists():
with open(self._FILE_PATH, "r") as f:
data = json.load(f)
bl = {k: v for k, v in data.get("blacklist", {}).items() if v > now}
rf = {k: v for k, v in data.get("refresh", {}).items()
if v.get("expires", 0) > now}
# 文件数据填充内存缺失项,不覆盖已在内存中手动设置的值
for k, v in bl.items():
self._fallback_blacklist.setdefault(k, v)
for k, v in rf.items():
self._fallback_refresh.setdefault(k, v)
except Exception:
pass
def _save_file(self):
"""将内存中的 token 数据写回 JSON 文件。"""
self._ensure_file_dir()
try:
data = {
"blacklist": self._fallback_blacklist,
"refresh": self._fallback_refresh,
"saved_at": time.time(),
}
# 原子写入:先写 .tmp 再 rename
tmp = self._FILE_PATH.with_suffix(".json.tmp")
with open(tmp, "w") as f:
json.dump(data, f)
tmp.replace(self._FILE_PATH)
except Exception:
pass
# ── Redis 连接 ──
async def _get_redis(self):
# 已有有效连接 → 快速返回
if self._redis and self._redis is not True:
try:
await self._redis.ping()
return self._redis
except Exception:
self._redis = None # 连接失效,下次重建
# Redis 不可用标记(None=未试过,False=已失败,不重试)
if self._redis is False:
return None
# 尝试连接(含超时,避免 Windows 下默认 4s 等待)
if self._redis is None or self._redis is True:
try:
from redis.asyncio import Redis
r = Redis.from_url(settings.REDIS_URL, decode_responses=True, socket_connect_timeout=1)
await r.ping()
self._redis = r
except Exception:
self._redis = False # 缓存失败,不再重试
return self._redis if self._redis and self._redis is not True else None
# ─── Access Token 黑名单(logout 时加入)───
async def blacklist(self, jti: str, ttl: int = 900):
r = await self._get_redis()
if r:
await r.setex(f"bl:{jti}", ttl, "1")
else:
expire = time.time() + ttl
self._fallback_blacklist[jti] = expire
self._save_file()
async def is_blacklisted(self, jti: str) -> bool:
r = await self._get_redis()
if r:
return await r.exists(f"bl:{jti}") > 0
self._load_file()
expire = self._fallback_blacklist.get(jti, 0)
if expire and expire < time.time():
del self._fallback_blacklist[jti]
self._save_file()
return False
return bool(expire)
# ─── Refresh Token 管理 ───
async def store_refresh(self, jti: str, user_id: str, tenant_id: str, ttl: int = 2592000):
r = await self._get_redis()
if r:
await r.setex(f"rt:{jti}", ttl, f"{user_id}|{tenant_id}")
await r.sadd(f"user_rt:{user_id}", jti)
await r.expire(f"user_rt:{user_id}", ttl)
else:
self._load_file()
self._fallback_refresh[jti] = {"user_id": user_id, "tenant_id": tenant_id, "expires": time.time() + ttl}
self._save_file()
async def validate_refresh(self, jti: str) -> dict | None:
r = await self._get_redis()
found_in_redis = None
if r:
val = await r.get(f"rt:{jti}")
if val:
found_in_redis = val
parts = val.split("|")
return {"user_id": parts[0], "tenant_id": parts[1], "jti": jti}
# Redis 未找到时回退到文件(迁移期:网络分裂时存的 token 只落在文件里)
self._load_file()
stored = self._fallback_refresh.get(jti)
if stored and stored["expires"] > time.time():
if r and not found_in_redis:
await r.setex(f"rt:{jti}", int(stored["expires"] - time.time()),
f"{stored['user_id']}|{stored.get('tenant_id','')}")
await r.sadd(f"user_rt:{stored['user_id']}", jti)
return {"user_id": stored["user_id"], "tenant_id": stored.get("tenant_id", ""), "jti": jti}
return None
async def revoke_refresh(self, jti: str):
r = await self._get_redis()
if r:
await r.delete(f"rt:{jti}")
else:
self._load_file()
self._fallback_refresh.pop(jti, None)
self._save_file()
async def revoke_all_refresh_for_user(self, user_id: str):
"""吊销指定用户的所有 refresh token"""
r = await self._get_redis()
if r:
jtis = await r.smembers(f"user_rt:{user_id}")
if jtis:
await r.delete(*[f"rt:{jti}" for jti in jtis])
await r.delete(f"user_rt:{user_id}")
else:
self._load_file()
to_del = [jti for jti, data in self._fallback_refresh.items()
if data.get("user_id") == user_id]
for jti in to_del:
self._fallback_refresh.pop(jti, None)
if to_del:
self._save_file()
# ─── Refresh Token 重放检测(被盗检测)───
async def mark_consumed(self, jti: str, ttl: int = 604800):
"""将已使用过的 refresh token jti 标记为 consumed,用于重放检测"""
r = await self._get_redis()
if r:
await r.setex(f"consumed:{jti}", ttl, "1")
else:
self._load_file()
self._fallback_refresh.pop(jti, None)
async def is_consumed(self, jti: str) -> bool:
"""检查 jti 是否已被消费过(重放攻击检测)"""
r = await self._get_redis()
if r:
return await r.exists(f"consumed:{jti}") > 0
return False
# 全局单例
token_store = TokenStore()