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.
198 lines
7.3 KiB
Python
198 lines
7.3 KiB
Python
"""JWT Token 持久存储 — Redis 黑名单 + Refresh Token 管理
|
||
|
||
三层 fallback:Redis → 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()
|