Files
stock/backend/app/auth.py
cirry 2f9c8bee2b 日K加载提速:adj_factor 覆盖索引 + 冷路径并发 + 热路径进程内缓存直返
- adj_factor 覆盖索引 (ts_code,trade_date) INCLUDE (adj_factor):根治堆碎片化
  (单股 6516 行散 6516 块,满载时位图堆扫 1.5s+),Index Only Scan ~2ms;
  因子查询全部改 2 列投影,lag() 窗口只取变点(6516→32 行)
- preview 冷路径 2 会话 2 波并发,信息卡合并为 LEFT JOIN LATERAL 一条
- 鉴权会话 60s 进程内缓存(登出/全端登出即时失效),全站请求省 ~80ms
- pvj/chipsj/stocksj/facetsj 存 model_dump_json 原串直返(与 response_model
  字节一致),热路径 230-2190ms → 1-2ms;get_version 本地缓存+bump 即时可见
- 连接池 10+20;smoke_test 适配鉴权缓存
2026-09-02 16:52:28 +08:00

138 lines
4.5 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.
"""密码校验、数据库会话与 FastAPI 鉴权依赖。"""
from __future__ import annotations
import hashlib
import secrets
import time
from datetime import datetime, timedelta, timezone
from argon2 import PasswordHasher
from argon2.exceptions import InvalidHashError, VerificationError
from fastapi import Cookie, Depends, HTTPException, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from .config import settings
from .db import get_session
from .models import AuthSession, User
password_hasher = PasswordHasher(
time_cost=2,
memory_cost=19_456,
parallelism=1,
hash_len=32,
salt_len=16,
)
# 不存在的用户名也执行一次 Argon2降低用户名枚举与计时攻击差异。
_DUMMY_HASH = password_hasher.hash("not-a-real-password")
def utcnow() -> datetime:
return datetime.now(timezone.utc)
def hash_password(password: str) -> str:
return password_hasher.hash(password)
def verify_password(password_hash: str, password: str) -> bool:
try:
return password_hasher.verify(password_hash, password)
except (VerificationError, InvalidHashError):
return False
def verify_dummy_password(password: str) -> None:
verify_password(_DUMMY_HASH, password)
def new_session_token() -> str:
return secrets.token_urlsafe(48)
def token_digest(token: str) -> str:
return hashlib.sha256(token.encode("utf-8")).hexdigest()
def session_expiry() -> datetime:
return utcnow() + timedelta(hours=settings.auth_session_hours)
def unauthorized(detail: str = "登录状态无效或已过期") -> HTTPException:
return HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=detail,
headers={"WWW-Authenticate": "Session"},
)
async def get_auth_session(
token: str,
db: AsyncSession,
) -> AuthSession | None:
now = utcnow()
stmt = (
select(AuthSession)
.options(joinedload(AuthSession.user))
.where(
AuthSession.token_hash == token_digest(token),
AuthSession.revoked_at.is_(None),
AuthSession.expires_at > now,
)
)
auth_session = (await db.execute(stmt)).scalar_one_or_none()
if auth_session is None or not auth_session.user.is_active:
return None
# 避免每个 API 请求都写数据库;最多每 5 分钟刷新一次活动时间。
if auth_session.last_seen_at < now - timedelta(minutes=5):
auth_session.last_seen_at = now
await db.commit()
return auth_session
# --- 鉴权会话进程内缓存 --------------------------------------------------------
# 原本每个 API 请求都要为鉴权付 ~3 个远程 RTT连接池 pre_ping + 查 AuthSession +
# joinedload User命中后完全跳过 DB。登出/撤销在 auth_api 同进程立即清;
# 其他进程撤销(多 worker 部署)最长 _SESSION_CACHE_TTL 后自然过期。
# 缓存的是脱管 ORM User属性已加载、expire_on_commit=False脱管访问安全
_SESSION_CACHE_TTL = 60.0
_session_cache: dict[str, tuple[float, datetime, User]] = {} # digest -> (mono 到期, 会话到期, user)
_SESSION_CACHE_MAX = 256
def drop_session_cache(digest: str | None = None, user_id: int | None = None) -> None:
"""登出/撤销时清缓存:按 token 或按用户logout-all"""
if digest is not None:
_session_cache.pop(digest, None)
return
if user_id is not None:
for k in [k for k, (_, _, u) in _session_cache.items() if u.id == user_id]:
_session_cache.pop(k)
async def require_user(
stock_session: str | None = Cookie(default=None, alias=settings.auth_cookie_name),
db: AsyncSession = Depends(get_session),
) -> User:
if not stock_session:
raise unauthorized()
d = token_digest(stock_session)
hit = _session_cache.get(d)
if hit is not None:
expires_mono, sess_expires, user = hit
if expires_mono > time.monotonic() and sess_expires > utcnow() and user.is_active:
return user
_session_cache.pop(d, None) # 过期/失效条目顺手清掉
auth_session = await get_auth_session(stock_session, db)
if auth_session is None:
raise unauthorized()
ttl = min(_SESSION_CACHE_TTL, max(1.0, (auth_session.expires_at - utcnow()).total_seconds()))
_session_cache[d] = (time.monotonic() + ttl, auth_session.expires_at, auth_session.user)
while len(_session_cache) > _SESSION_CACHE_MAX:
_session_cache.pop(next(iter(_session_cache)))
return auth_session.user