日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 适配鉴权缓存
This commit is contained in:
49
backend/alembic/versions/20260902_01_adj_factor_cover.py
Normal file
49
backend/alembic/versions/20260902_01_adj_factor_cover.py
Normal file
@@ -0,0 +1,49 @@
|
||||
"""adj_factor 覆盖索引:因子查询免堆访问(碎片免疫)
|
||||
|
||||
Revision ID: 20260902_01
|
||||
Revises: 20260815_04
|
||||
Create Date: 2026-09-02
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision: str = "20260902_01"
|
||||
down_revision: Union[str, Sequence[str], None] = "20260815_04"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# CONCURRENTLY / VACUUM 不能在事务内执行;autocommit_block 内每条语句独立提交
|
||||
with op.get_context().autocommit_block():
|
||||
op.create_index(
|
||||
"ix_adj_code_date_cover",
|
||||
"adj_factor",
|
||||
["ts_code", "trade_date"],
|
||||
postgresql_concurrently=True,
|
||||
postgresql_include=["adj_factor"],
|
||||
)
|
||||
# ts_code 单列索引被覆盖索引取代(同前缀),删掉省写放大
|
||||
op.drop_index(
|
||||
"ix_adj_factor_ts_code",
|
||||
table_name="adj_factor",
|
||||
postgresql_concurrently=True,
|
||||
)
|
||||
# 刷可见性映射(Index Only Scan 的前提)+ 更新统计信息;低峰执行(3.2GB 表数十秒)
|
||||
op.execute("VACUUM (ANALYZE) adj_factor")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.get_context().autocommit_block():
|
||||
op.create_index(
|
||||
"ix_adj_factor_ts_code",
|
||||
"adj_factor",
|
||||
["ts_code"],
|
||||
postgresql_concurrently=True,
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_adj_code_date_cover",
|
||||
table_name="adj_factor",
|
||||
postgresql_concurrently=True,
|
||||
)
|
||||
@@ -18,7 +18,7 @@ import json
|
||||
from datetime import datetime
|
||||
|
||||
import pandas as pd
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Response, UploadFile
|
||||
from sqlalchemy import delete, func, select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.sql.elements import TextClause
|
||||
@@ -33,7 +33,7 @@ from .data import fetcher, repository, tushare_provider
|
||||
from .data.aggregation import bars_per_year, resample_bars
|
||||
from .data.market_overview import MarketOverviewError, fetch_overview
|
||||
from .data.symbols import plain_code
|
||||
from .db import get_session
|
||||
from .db import async_session, get_session
|
||||
from .domain import Bar
|
||||
from . import indicators as ind
|
||||
from .trades import parse_statement
|
||||
@@ -41,7 +41,6 @@ from .models import (
|
||||
AdjFactor,
|
||||
BacktestRun,
|
||||
Candle,
|
||||
DailySnapshot,
|
||||
ScreenerQuery,
|
||||
StockBasic,
|
||||
UserPreference,
|
||||
@@ -89,6 +88,25 @@ from .screener.llm import ScreenerError, parse_conditions, parse_event_spec
|
||||
router = APIRouter(prefix="/api", dependencies=[Depends(require_user)])
|
||||
|
||||
|
||||
def _raw_json(resp) -> str:
|
||||
"""pydantic-core(Rust)序列化:与 response_model 直返时的字节完全一致(紧凑分隔符、
|
||||
非 ASCII 直出、浮点小数形式),且比 stdlib json.dumps 快。大响应(preview ~250KB)
|
||||
命中缓存时直接 Response 原样返回,跳过校验/再序列化。"""
|
||||
return resp.model_dump_json()
|
||||
|
||||
|
||||
async def _cached_json_response(key: str) -> Response | None:
|
||||
"""两级缓存读(进程内 → Redis):命中返回可直接吐给客户端的 Response。
|
||||
存的均为序列化好的 JSON 字符串(Redis 侧 json.loads 后仍是 str),Redis 命中顺手晋级本地。"""
|
||||
raw = cache.local_get(key)
|
||||
if raw is None:
|
||||
raw = await cache.cache_get(key)
|
||||
if not isinstance(raw, str):
|
||||
return None
|
||||
cache.local_set(key, raw, ttl=120)
|
||||
return Response(content=raw, media_type="application/json")
|
||||
|
||||
|
||||
def _series_to_jsonable(s: pd.Series) -> list[float | None]:
|
||||
"""NaN -> None(lightweight-charts 的 whitespace data,跳过指标预热期)。"""
|
||||
out: list[float | None] = []
|
||||
@@ -114,6 +132,40 @@ _ADJUST_MODES = ("bfq", "qfq", "hfq")
|
||||
# MA 全量集合(前端已改为本地计算 MA,后端始终返回此集合以保证缓存一致)
|
||||
_FULL_MA_SET = (5, 10, 20, 30, 60, 120, 250)
|
||||
|
||||
# 信息卡一条 SQL 拿全:stock_basic 基本信息 + 「优先与行情同日、缺则最新日」的 daily_snapshot
|
||||
# (LATERAL 单条替换原两条查询,语义不变:target 为 NULL 时全按最新日兜底)
|
||||
_INFO_SQL = text(
|
||||
"""
|
||||
SELECT sb.ts_code, sb.symbol, sb.name, sb.industry, sb.area, sb.market, sb.list_date,
|
||||
ds.turnover_rate, ds.pe_ttm, ds.pb, ds.total_mv, ds.circ_mv
|
||||
FROM stock_basic sb
|
||||
LEFT JOIN LATERAL (
|
||||
SELECT turnover_rate, pe_ttm, pb, total_mv, circ_mv
|
||||
FROM daily_snapshot
|
||||
WHERE ts_code = sb.ts_code
|
||||
ORDER BY (trade_date = cast(:target AS timestamp)) DESC, trade_date DESC
|
||||
LIMIT 1
|
||||
) ds ON true
|
||||
WHERE sb.ts_code = :code
|
||||
"""
|
||||
)
|
||||
|
||||
# 复权因子是阶梯函数(除权日之间不变):只取「变化点」行,把每符号 ~7000 行日级因子压到
|
||||
# 几十行(600118 仅 32 行),传输量再降两个数量级;lag 窗口在覆盖索引上走 Index Only Scan。
|
||||
# upto 传全局最新交易日(非分页)或 end 日期(分页);窗口首行 prev 为 NULL 恒被保留(窗口基线因子)。
|
||||
_FACTOR_STEP_SQL = text(
|
||||
"""
|
||||
SELECT trade_date, adj_factor FROM (
|
||||
SELECT trade_date, adj_factor,
|
||||
lag(adj_factor) OVER (ORDER BY trade_date) AS prev
|
||||
FROM adj_factor
|
||||
WHERE ts_code = :code AND trade_date <= :upto
|
||||
) t
|
||||
WHERE adj_factor IS DISTINCT FROM prev
|
||||
ORDER BY trade_date
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _adjust_bars(bars: list[Bar], factors, from_mode: str, to_mode: str) -> list[Bar]:
|
||||
"""按复权因子把 K 线从 from_mode 换算到 to_mode(bfq/qfq/hfq)。
|
||||
@@ -121,7 +173,7 @@ def _adjust_bars(bars: list[Bar], factors, from_mode: str, to_mode: str) -> list
|
||||
相对不复权的乘数:bfq=1,qfq=f(t)/f(latest),hfq=f(t)。
|
||||
因子缺失的日期向前沿用最近因子(因子是阶梯函数,除权日之间不变)。
|
||||
"""
|
||||
fd = sorted((f.trade_date.date(), float(f.adj_factor)) for f in factors)
|
||||
fd = sorted((f[0].date(), float(f[1])) for f in factors)
|
||||
fdates = [d for d, _ in fd]
|
||||
f_latest = fd[-1][1]
|
||||
|
||||
@@ -262,26 +314,27 @@ async def list_stocks(
|
||||
offset: int = 0,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
user=Depends(require_user),
|
||||
) -> StockListResponse:
|
||||
) -> Response:
|
||||
"""全市场股票列表:stock_basic 基本信息 + candles 最新行情 + daily_snapshot 估值指标
|
||||
(换手率/PE-TTM/PB/市值,无快照则这些列为空)。
|
||||
watched_only=true 只看自选(自选有独立的「自选」分类入口,列表不再把自选排最前)。
|
||||
sort ∈ {symbol,total_mv,circ_mv,pe_ttm,pb,turnover_rate}(白名单,其他值回落 symbol),
|
||||
order ∈ asc/desc;快照列排序时缺失值(无快照/亏损无 PE)恒排末尾。
|
||||
Redis 缓存:按「用户自选版本 + 查询参数(含排序)」缓存整页(含 total);自选增删即时失效。"""
|
||||
缓存:按「用户自选版本 + 查询参数(含排序)」缓存整页(含 total);自选增删即时失效;
|
||||
与 preview 同款序列化 JSON 直返(j 前缀),命中跳过 pydantic 校验/序列化。"""
|
||||
search = search.strip()
|
||||
sort = sort if sort in _STOCKS_SORTS else "symbol"
|
||||
order = "desc" if order.lower() == "desc" else "asc"
|
||||
limit = max(1, min(limit, 500))
|
||||
offset = max(0, offset)
|
||||
key = (
|
||||
f"stocks:u{user.id}"
|
||||
f"stocksj:u{user.id}"
|
||||
f":v{await cache.get_version(f'watchlist:{user.id}')}"
|
||||
f":{cache.digest(search, market, industry, area, watched_only, sort, order, limit, offset)}"
|
||||
)
|
||||
cached = await cache.cache_get(key)
|
||||
cached = await _cached_json_response(key)
|
||||
if cached is not None:
|
||||
return StockListResponse(**cached)
|
||||
return cached
|
||||
params = {
|
||||
"search": search,
|
||||
"psearch": f"%{search}%",
|
||||
@@ -296,16 +349,18 @@ async def list_stocks(
|
||||
total = (await session.execute(_STOCKS_COUNT_SQL, params)).scalar_one()
|
||||
rows = (await session.execute(_stocks_sql(sort, order), params)).mappings().all()
|
||||
resp = StockListResponse(total=total, items=[StockListItemOut(**r) for r in rows])
|
||||
await cache.cache_set(key, resp.model_dump(mode="json"), settings.stocks_cache_ttl)
|
||||
return resp
|
||||
raw = _raw_json(resp)
|
||||
cache.local_set(key, raw, ttl=min(120, settings.stocks_cache_ttl))
|
||||
asyncio.create_task(cache.cache_set(key, raw, ttl=settings.stocks_cache_ttl))
|
||||
return Response(content=raw, media_type="application/json")
|
||||
|
||||
|
||||
@router.get("/stocks/facets", response_model=StockFacetsResponse)
|
||||
async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFacetsResponse:
|
||||
async def stock_facets(session: AsyncSession = Depends(get_session)) -> Response:
|
||||
"""看股页筛选项:行业 / 地域(含数量,按数量降序)。stock_basic 很少变,长缓存。"""
|
||||
cached = await cache.cache_get("facets:stocks")
|
||||
cached = await _cached_json_response("facetsj:stocks")
|
||||
if cached is not None:
|
||||
return StockFacetsResponse(**cached)
|
||||
return cached
|
||||
industries = (
|
||||
await session.execute(text("""
|
||||
SELECT industry AS name, count(*) AS n FROM stock_basic
|
||||
@@ -324,8 +379,10 @@ async def stock_facets(session: AsyncSession = Depends(get_session)) -> StockFac
|
||||
industries=[FacetItemOut(name=r["name"], count=r["n"]) for r in industries],
|
||||
areas=[FacetItemOut(name=r["name"], count=r["n"]) for r in areas],
|
||||
)
|
||||
await cache.cache_set("facets:stocks", resp.model_dump(mode="json"), settings.facets_cache_ttl)
|
||||
return resp
|
||||
raw = _raw_json(resp)
|
||||
cache.local_set("facetsj:stocks", raw, ttl=min(120, settings.facets_cache_ttl))
|
||||
asyncio.create_task(cache.cache_set("facetsj:stocks", raw, ttl=settings.facets_cache_ttl))
|
||||
return Response(content=raw, media_type="application/json")
|
||||
|
||||
|
||||
@router.get("/market/overview", response_model=MarketOverviewResponse)
|
||||
@@ -788,7 +845,7 @@ async def screener_preview(
|
||||
ts_code: str, limit: int = 500, adjust: str = "qfq", timeframe: str = "1d", mas: str = "5,10,20,60",
|
||||
zx: str = "10,20,30,60", end: str | None = None,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> PreviewResponse:
|
||||
) -> Response:
|
||||
"""个股详情预览:日线(candles 不复权底座 + adj_factor 本地换算 bfq/qfq/hfq,
|
||||
未缓存自动拉取,落后全市场最新交易日则强制刷新)+ 全套指标 + 最新截面信息卡。
|
||||
timeframe 聚合到周/月/年(先复权再聚合);mas 指定主图 MA 周期(逗号分隔)。
|
||||
@@ -819,73 +876,92 @@ async def screener_preview(
|
||||
raise HTTPException(status_code=400, detail="end 格式应为 YYYY-MM-DD")
|
||||
symbol = plain_code(ts_code)
|
||||
|
||||
# --- Redis 读缓存:历史窗口(end 翻页)只增不改,最新窗口每日由全市场同步推进;
|
||||
# 键含 ver:candles 版本号(同步完成后自增,旧缓存全部失效),TTL 兜底(cache.py)
|
||||
# --- 两级读缓存:历史窗口(end 翻页)只增不改,最新窗口每日由全市场同步推进;
|
||||
# 键含 ver:candles 版本号(同步完成后自增,旧缓存全部失效),TTL 兜底(cache.py)。
|
||||
# 存序列化好的 JSON 直返(j: 前缀),跳过 json.loads + pydantic 校验/序列化(热路径数百 ms → 个位数)。
|
||||
# 注:ma_periods 不参与缓存键 —— 前端已改为本地计算 MA,后端始终返回全量 MA 集合
|
||||
cache_key = cache.digest(
|
||||
"preview", ts_code, timeframe, limit, adjust,
|
||||
end_dt.strftime("%Y-%m-%d") if end_dt else None,
|
||||
await cache.get_version("candles"),
|
||||
)
|
||||
cached = await cache.cache_get(f"pv:{cache_key}")
|
||||
cached = await _cached_json_response(f"pvj:{cache_key}")
|
||||
if cached is not None:
|
||||
return PreviewResponse.model_validate(cached)
|
||||
return cached
|
||||
|
||||
# --- 日线:candles(全量不复权底座);未缓存拉取,落后于全市场最新交易日则强制刷新 ---
|
||||
# fetcher 只做「不复权」增量 upsert,底座口径恒为 bfq(TDX 全量 + Tushare 增量),
|
||||
# 复权(qfq/hfq)读取时按 adj_factor 表本地换算。
|
||||
# 每次只取「窗口 + 400 根预热」行(MA250/MACD EMA 在 400 根内充分收敛),不拉全量:
|
||||
# 首屏 ~500 根秒开,向左滚动时按 end 参数逐页向前翻。
|
||||
global_latest = await session.scalar(
|
||||
select(func.max(Candle.ts)).where(Candle.timeframe == "1d")
|
||||
)
|
||||
frame_mult = {"1d": 1, "1w": 6, "1M": 24, "1y": 280}[timeframe]
|
||||
fetch_n = min(100000, limit * frame_mult + 400)
|
||||
source = "bfq"
|
||||
mode = "bfq"
|
||||
if end_dt is not None:
|
||||
# 向前翻页:取 end 之前的历史窗口,不触发同步(历史浏览)
|
||||
rows = await repository.get_candles_before(session, symbol, "1d", before=end_dt, limit=fetch_n)
|
||||
else:
|
||||
# 注意取「最新 fetch_n 根」而非最旧:get_candles 是 asc+limit(取最旧),窗口化后首屏会停在过期日期
|
||||
rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n)
|
||||
try:
|
||||
if not rows:
|
||||
await fetcher.sync_symbol(session, symbol, source="auto")
|
||||
rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n)
|
||||
elif global_latest is not None and rows[-1].ts.date() < global_latest.date():
|
||||
await fetcher.sync_symbol(session, symbol, source="auto", force=True)
|
||||
rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n)
|
||||
except Exception: # noqa: BLE001 —— tushare/写库失败时回滚会话(否则毒化后兜底查询 500)
|
||||
await session.rollback()
|
||||
if not rows:
|
||||
rows = []
|
||||
bars = _rows_to_bars(rows)
|
||||
# 并发约定:注入 session 与 s2 各占一条连接,每次 gather 里每个 session 恰好跑一条查询
|
||||
# (AsyncSession 单连接非并发安全),把 ~6 次串行 DB RTT 折叠成 2 个波次。
|
||||
async with async_session() as s2:
|
||||
if end_dt is not None:
|
||||
# 向前翻页:取 end 之前的历史窗口,不触发同步(历史浏览);max(ts) 用不到
|
||||
rows = await repository.get_candles_before(session, symbol, "1d", before=end_dt, limit=fetch_n)
|
||||
global_latest = None
|
||||
else:
|
||||
# Wave 1:candles 窗口(注入 session)+ 全市场最新交易日(s2)并行
|
||||
rows, global_latest = await asyncio.gather(
|
||||
repository.get_recent_candles(session, symbol, "1d", limit=fetch_n),
|
||||
s2.scalar(select(func.max(Candle.ts)).where(Candle.timeframe == "1d")),
|
||||
)
|
||||
try:
|
||||
if not rows:
|
||||
await fetcher.sync_symbol(session, symbol, source="auto")
|
||||
rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n)
|
||||
elif global_latest is not None and rows[-1].ts.date() < global_latest.date():
|
||||
await fetcher.sync_symbol(session, symbol, source="auto", force=True)
|
||||
rows = await repository.get_recent_candles(session, symbol, "1d", limit=fetch_n)
|
||||
except Exception: # noqa: BLE001 —— tushare/写库失败时回滚会话(否则毒化后兜底查询 500)
|
||||
await session.rollback()
|
||||
if not rows:
|
||||
rows = []
|
||||
bars = _rows_to_bars(rows)
|
||||
|
||||
if not bars and end_dt is None:
|
||||
raise HTTPException(status_code=404, detail=f"无数据: {ts_code}(可先点「同步市场数据」)")
|
||||
# 信息卡取未聚合的日线最新 bar(聚合后 ts 是周期起点,不适用于「最新交易日」)
|
||||
last_daily = bars[-1] if bars else None
|
||||
prev_daily = bars[-2] if len(bars) > 1 else None
|
||||
# 翻页到底(end 之前无数据):返回空页 + has_more=False,前端停止向前翻页
|
||||
if not bars and end_dt is None:
|
||||
raise HTTPException(status_code=404, detail=f"无数据: {ts_code}(可先点「同步市场数据」)")
|
||||
# 信息卡取未聚合的日线最新 bar(聚合后 ts 是周期起点,不适用于「最新交易日」)
|
||||
last_daily = bars[-1] if bars else None
|
||||
prev_daily = bars[-2] if len(bars) > 1 else None
|
||||
# 翻页到底(end 之前无数据):返回空页 + has_more=False,前端停止向前翻页
|
||||
|
||||
# --- 复权换算:请求模式与底座模式不同时按 adj_factor 本地换算(无因子则维持原样) ---
|
||||
# 只取窗口内因子(qfq 归一还需全局最新因子,追加到最后一行即可,_adjust_bars 取 f_latest=末项)
|
||||
if adjust != mode and bars:
|
||||
fq = select(AdjFactor).where(AdjFactor.ts_code == ts_code)
|
||||
window_end = end_dt if end_dt is not None else bars[-1].ts
|
||||
if window_end is not None:
|
||||
fq = fq.where(AdjFactor.trade_date <= window_end)
|
||||
factors = list((await session.execute(fq.order_by(AdjFactor.trade_date))).scalars().all())
|
||||
# --- Wave 2:复权因子(s2,覆盖索引 Index Only Scan)+ 信息卡(注入 session,LATERAL 一条)并行 ---
|
||||
async def _fetch_factors() -> list | None:
|
||||
if adjust == mode or not bars:
|
||||
return None
|
||||
# 只取因子「变化点」行(覆盖索引 Index Only Scan,免堆访问——adj_factor 堆碎片化
|
||||
# 严重);bisect 在阶梯函数上取值与日级序列逐字节一致
|
||||
if end_dt is not None:
|
||||
# 分页:窗口 ≤ end 的变化点 + 全局最新因子(qfq 以最新因子归一)
|
||||
win = list((await s2.execute(_FACTOR_STEP_SQL, {"code": ts_code, "upto": end_dt})).all())
|
||||
if win:
|
||||
latest_f = (await s2.execute(
|
||||
select(AdjFactor.trade_date, AdjFactor.adj_factor)
|
||||
.where(AdjFactor.ts_code == ts_code)
|
||||
.order_by(AdjFactor.trade_date.desc()).limit(1)
|
||||
)).first()
|
||||
if latest_f is not None:
|
||||
win.append(latest_f)
|
||||
return win or None
|
||||
# 非分页:上界 global_latest(≥ 最新 bar),末项变化点即全局最新因子,比
|
||||
# 「窗口 ≤ bars[-1].ts + 单独 latest」少一次查询
|
||||
return list((await s2.execute(
|
||||
_FACTOR_STEP_SQL, {"code": ts_code, "upto": global_latest}
|
||||
)).all()) or None
|
||||
|
||||
factors, info_row = await asyncio.gather(
|
||||
_fetch_factors(),
|
||||
session.execute(_INFO_SQL, {"code": ts_code, "target": last_daily.ts if last_daily else None}),
|
||||
)
|
||||
|
||||
# --- 复权换算:请求模式与底座模式不同时按 adj_factor 本地换算(无因子则维持原样) ---
|
||||
if factors:
|
||||
latest_f = (
|
||||
await session.execute(
|
||||
select(AdjFactor).where(AdjFactor.ts_code == ts_code)
|
||||
.order_by(AdjFactor.trade_date.desc()).limit(1)
|
||||
)
|
||||
).scalars().first()
|
||||
if latest_f is not None:
|
||||
factors.append(latest_f)
|
||||
bars = _adjust_bars(bars, factors, mode, adjust)
|
||||
mode = adjust
|
||||
source = adjust
|
||||
@@ -927,24 +1003,19 @@ async def screener_preview(
|
||||
for key in group:
|
||||
group[key] = group[key][-limit:]
|
||||
|
||||
# --- 信息卡:stock_basic + candles 最新日线 bar + 与其对齐的快照(避免混用不同交易日) ---
|
||||
sb = (await session.execute(select(StockBasic).where(StockBasic.ts_code == ts_code))).scalars().first()
|
||||
ds = None
|
||||
if last_daily is not None:
|
||||
# 优先取与行情同日的快照;缺当日快照时退最新(字段可能与行情差日期,罕见)
|
||||
ds = (
|
||||
await session.execute(
|
||||
select(DailySnapshot).where(
|
||||
DailySnapshot.ts_code == ts_code, DailySnapshot.trade_date == last_daily.ts
|
||||
)
|
||||
)
|
||||
).scalars().first()
|
||||
if ds is None:
|
||||
ds = (
|
||||
await session.execute(
|
||||
select(DailySnapshot).where(DailySnapshot.ts_code == ts_code).order_by(DailySnapshot.trade_date.desc()).limit(1)
|
||||
)
|
||||
).scalars().first()
|
||||
# --- 信息卡:Wave 2 已并行取回(stock_basic + 与行情同日对齐的快照、缺则最新日,见 _INFO_SQL) ---
|
||||
row = info_row.first()
|
||||
if row is not None:
|
||||
m = row._mapping
|
||||
sb_name, sb_industry, sb_area, sb_market, sb_list_date = (
|
||||
m["name"], m["industry"], m["area"], m["market"], m["list_date"]
|
||||
)
|
||||
ds_turnover, ds_pe, ds_pb, ds_tmv, ds_cmv = (
|
||||
m["turnover_rate"], m["pe_ttm"], m["pb"], m["total_mv"], m["circ_mv"]
|
||||
)
|
||||
else:
|
||||
sb_name = sb_industry = sb_area = sb_market = sb_list_date = None
|
||||
ds_turnover = ds_pe = ds_pb = ds_tmv = ds_cmv = None
|
||||
|
||||
def _yi(v) -> float | None:
|
||||
if v is None:
|
||||
@@ -955,11 +1026,11 @@ async def screener_preview(
|
||||
info = PreviewInfoOut(
|
||||
ts_code=ts_code,
|
||||
symbol=symbol,
|
||||
name=sb.name if sb else ts_code,
|
||||
industry=sb.industry if sb else None,
|
||||
area=sb.area if sb else None,
|
||||
market=sb.market if sb else None,
|
||||
list_date=sb.list_date if sb else None,
|
||||
name=sb_name or ts_code,
|
||||
industry=sb_industry,
|
||||
area=sb_area,
|
||||
market=sb_market,
|
||||
list_date=sb_list_date,
|
||||
trade_date=last_daily.ts if last_daily else None,
|
||||
open=last_daily.open if last_daily else None,
|
||||
high=last_daily.high if last_daily else None,
|
||||
@@ -970,11 +1041,11 @@ async def screener_preview(
|
||||
if last_daily and prev_daily and prev_daily.close else None,
|
||||
volume_hand=round(last_daily.volume / 100, 0) if last_daily else None, # 股 -> 手
|
||||
amount_yi=round(last_daily.amount / 1e8, 2) if last_daily and last_daily.amount else None, # 元 -> 亿元
|
||||
turnover_rate=ds.turnover_rate if ds else None,
|
||||
pe_ttm=ds.pe_ttm if ds else None,
|
||||
pb=ds.pb if ds else None,
|
||||
total_mv=_yi(ds.total_mv) if ds else None,
|
||||
circ_mv=_yi(ds.circ_mv) if ds else None,
|
||||
turnover_rate=ds_turnover,
|
||||
pe_ttm=ds_pe,
|
||||
pb=ds_pb,
|
||||
total_mv=_yi(ds_tmv),
|
||||
circ_mv=_yi(ds_cmv),
|
||||
)
|
||||
|
||||
candles = [
|
||||
@@ -984,15 +1055,18 @@ async def screener_preview(
|
||||
]
|
||||
resp = PreviewResponse(ts_code=ts_code, symbol=symbol, source=source, info=info,
|
||||
candles=candles, indicators=indicators, has_more=has_more)
|
||||
await cache.cache_set(f"pv:{cache_key}", resp.model_dump(mode="json"), ttl=600)
|
||||
return resp
|
||||
# 只序列化一次:本地(同步,120s)+ Redis(后台写,600s TTL 兜底跨进程/重启)
|
||||
raw = _raw_json(resp)
|
||||
cache.local_set(f"pvj:{cache_key}", raw, ttl=120)
|
||||
asyncio.create_task(cache.cache_set(f"pvj:{cache_key}", raw, ttl=600))
|
||||
return Response(content=raw, media_type="application/json")
|
||||
|
||||
|
||||
@router.get("/stock/{ts_code}/chips", response_model=ChipsResponse)
|
||||
async def stock_chips(
|
||||
ts_code: str, date: str | None = None, adjust: str = "qfq",
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> ChipsResponse:
|
||||
) -> Response:
|
||||
"""个股筹码峰(Tushare cyq_chips + cyq_perf,数据自 2018 年起)。
|
||||
|
||||
date=YYYY-MM-DD 为参考日(日 K 传当日;周/月 K 由前端传周期末):返回
|
||||
@@ -1008,11 +1082,12 @@ async def stock_chips(
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail="date 格式应为 YYYY-MM-DD")
|
||||
|
||||
# 历史截面不可变;键带 candles 版本号(adj_factor 随同步更新后旧缓存失效)
|
||||
# 历史截面不可变;键带 candles 版本号(adj_factor 随同步更新后旧缓存失效)。
|
||||
# 同 preview:存序列化 JSON 直返(j: 前缀),跳过 pydantic 校验/序列化。
|
||||
cache_key = cache.digest("chips", ts_code, ref or "latest", adjust, await cache.get_version("candles"))
|
||||
cached = await cache.cache_get(f"chips:{cache_key}")
|
||||
cached = await _cached_json_response(f"chipsj:{cache_key}")
|
||||
if cached is not None:
|
||||
return ChipsResponse.model_validate(cached)
|
||||
return cached
|
||||
|
||||
try:
|
||||
perf, rows = await asyncio.to_thread(tushare_provider.fetch_chips, ts_code, ref)
|
||||
@@ -1023,8 +1098,10 @@ async def stock_chips(
|
||||
ts_code=ts_code, trade_date=None, adjust=adjust,
|
||||
error="无筹码数据(cyq 数据自 2018 年起,或参考日早于数据起点)",
|
||||
)
|
||||
await cache.cache_set(f"chips:{cache_key}", resp.model_dump(mode="json"), ttl=3600)
|
||||
return resp
|
||||
raw = _raw_json(resp)
|
||||
cache.local_set(f"chipsj:{cache_key}", raw, ttl=120)
|
||||
await cache.cache_set(f"chipsj:{cache_key}", raw, ttl=3600)
|
||||
return Response(content=raw, media_type="application/json")
|
||||
|
||||
d = datetime.strptime(str(perf["trade_date"]), "%Y%m%d")
|
||||
|
||||
@@ -1032,16 +1109,17 @@ async def stock_chips(
|
||||
mult = 1.0
|
||||
if adjust != "bfq":
|
||||
factors = (await session.execute(
|
||||
select(AdjFactor).where(AdjFactor.ts_code == ts_code, AdjFactor.trade_date <= d)
|
||||
select(AdjFactor.trade_date, AdjFactor.adj_factor)
|
||||
.where(AdjFactor.ts_code == ts_code, AdjFactor.trade_date <= d)
|
||||
.order_by(AdjFactor.trade_date)
|
||||
)).scalars().all()
|
||||
)).all()
|
||||
if factors:
|
||||
latest_f = (await session.execute(
|
||||
select(AdjFactor).where(AdjFactor.ts_code == ts_code)
|
||||
select(AdjFactor.trade_date, AdjFactor.adj_factor).where(AdjFactor.ts_code == ts_code)
|
||||
.order_by(AdjFactor.trade_date.desc()).limit(1)
|
||||
)).scalars().first()
|
||||
f_at = float(factors[-1].adj_factor) # <=d 的最近因子(因子是阶梯函数)
|
||||
f_latest = float(latest_f.adj_factor) if latest_f else f_at
|
||||
)).first()
|
||||
f_at = float(factors[-1][1]) # <=d 的最近因子(因子是阶梯函数)
|
||||
f_latest = float(latest_f[1]) if latest_f else f_at
|
||||
mult = f_at / f_latest if adjust == "qfq" else f_at
|
||||
|
||||
def _px(v) -> float | None:
|
||||
@@ -1059,5 +1137,7 @@ async def stock_chips(
|
||||
weight_avg=_px(perf.get("weight_avg")),
|
||||
winner_rate=_px(perf.get("winner_rate")),
|
||||
)
|
||||
await cache.cache_set(f"chips:{cache_key}", resp.model_dump(mode="json"), ttl=21600)
|
||||
return resp
|
||||
raw = _raw_json(resp)
|
||||
cache.local_set(f"chipsj:{cache_key}", raw, ttl=120)
|
||||
await cache.cache_set(f"chipsj:{cache_key}", raw, ttl=21600)
|
||||
return Response(content=raw, media_type="application/json")
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from argon2 import PasswordHasher
|
||||
@@ -10,7 +11,7 @@ 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 selectinload
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from .config import settings
|
||||
from .db import get_session
|
||||
@@ -75,7 +76,7 @@ async def get_auth_session(
|
||||
now = utcnow()
|
||||
stmt = (
|
||||
select(AuthSession)
|
||||
.options(selectinload(AuthSession.user))
|
||||
.options(joinedload(AuthSession.user))
|
||||
.where(
|
||||
AuthSession.token_hash == token_digest(token),
|
||||
AuthSession.revoked_at.is_(None),
|
||||
@@ -93,13 +94,44 @@ async def get_auth_session(
|
||||
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
|
||||
|
||||
@@ -8,6 +8,7 @@ from sqlalchemy import delete, select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from .auth import (
|
||||
drop_session_cache,
|
||||
get_auth_session,
|
||||
new_session_token,
|
||||
require_user,
|
||||
@@ -131,6 +132,7 @@ async def logout(
|
||||
.values(revoked_at=utcnow())
|
||||
)
|
||||
await db.commit()
|
||||
drop_session_cache(digest=token_digest(stock_session)) # 同进程立即失效,登出即时生效
|
||||
clear_session_cookie(response)
|
||||
response.status_code = status.HTTP_204_NO_CONTENT
|
||||
response.headers["Cache-Control"] = "no-store"
|
||||
@@ -149,6 +151,7 @@ async def logout_all(
|
||||
.values(revoked_at=utcnow())
|
||||
)
|
||||
await db.commit()
|
||||
drop_session_cache(user_id=user.id) # 该用户全部 token 的本地缓存立即失效
|
||||
clear_session_cookie(response)
|
||||
response.status_code = status.HTTP_204_NO_CONTENT
|
||||
response.headers["Cache-Control"] = "no-store"
|
||||
|
||||
@@ -6,11 +6,16 @@
|
||||
旧缓存 key 里带着旧版本号,无需 SCAN 批量删除。
|
||||
- 只缓存「读多写少、可容忍短暂陈旧」的聚合数据(股票列表、筛选项等);
|
||||
K线/回测等口径敏感数据不走这里。
|
||||
- Redis 之前还有一层进程内本地缓存(local_get/local_set,0 RTT):只存已序列化好的
|
||||
JSON 字符串,命中后接口直接 Response(content=raw) 原样返回,跳过 json.loads +
|
||||
pydantic 校验/序列化(大响应这两步合计可达数百 ms)。本地 TTL 恒 ≤ Redis TTL,
|
||||
多进程部署时本地条目最多比 Redis 多陈旧 120s;版本号 bump 在同进程立即生效。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
@@ -48,6 +53,46 @@ def digest(*parts: Any) -> str:
|
||||
return hashlib.md5(raw.encode()).hexdigest() # noqa: S324
|
||||
|
||||
|
||||
# --- 进程内本地缓存层(Redis 之前,0 RTT)-----------------------------------
|
||||
# asyncio 单线程,dict 读写无锁安全;按插入序淘汰近似 LRU(命中不续期,足够用)。
|
||||
_LOCAL_MAX_BYTES = 64 * 1024 * 1024 # 总字节预算(preview 响应 ~250KB,留足 1y 大窗口)
|
||||
_LOCAL_MAX_ENTRIES = 128
|
||||
_local_store: dict[str, tuple[float, str]] = {} # key -> (expires_monotonic, raw_json)
|
||||
_local_bytes = 0
|
||||
|
||||
|
||||
def local_get(key: str) -> str | None:
|
||||
"""命中返回已序列化的 JSON 字符串(调用方直接 Response 原样返回)。"""
|
||||
global _local_bytes
|
||||
ent = _local_store.get(key)
|
||||
if ent is None:
|
||||
return None
|
||||
expires, raw = ent
|
||||
if expires < time.monotonic():
|
||||
_local_store.pop(key, None)
|
||||
_local_bytes -= len(raw) + 64
|
||||
return None
|
||||
return raw
|
||||
|
||||
|
||||
def local_set(key: str, raw: str, ttl: int) -> None:
|
||||
"""ttl 为秒;调用方应传 min(本地默认, Redis TTL),保证本地不比 Redis 活得久。"""
|
||||
global _local_bytes
|
||||
old = _local_store.pop(key, None)
|
||||
if old is not None:
|
||||
_local_bytes -= len(old[1])
|
||||
_local_store[key] = (time.monotonic() + max(1, min(ttl, 120)), raw)
|
||||
_local_bytes += len(raw) + 64 # 连同 dict/tuple 开销粗略计入
|
||||
while _local_store and (len(_local_store) > _LOCAL_MAX_ENTRIES or _local_bytes > _LOCAL_MAX_BYTES):
|
||||
_local_bytes -= len(_local_store.popitem(last=False)[1][1]) + 64
|
||||
|
||||
|
||||
# --- 版本号本地缓存:热请求连 Redis GET ver:xx 都省掉 ------------------------
|
||||
# 同进程 bump_version 立即刷新本地;其他进程 bump 后本地最多陈旧 _LOCAL_VERSION_TTL。
|
||||
_LOCAL_VERSION_TTL = 60.0
|
||||
_local_versions: dict[str, tuple[int, float]] = {} # name -> (version, fetched_monotonic)
|
||||
|
||||
|
||||
async def cache_get(key: str) -> Any | None:
|
||||
c = _client()
|
||||
if c is None:
|
||||
@@ -72,13 +117,20 @@ async def cache_set(key: str, value: Any, ttl: int) -> None:
|
||||
|
||||
|
||||
async def get_version(name: str) -> int:
|
||||
"""读版本号(缺省 0)。版本号参与缓存 key:INCR 后旧 key 全部失效。"""
|
||||
"""读版本号(缺省 0)。版本号参与缓存 key:INCR 后旧 key 全部失效。
|
||||
先查本地(60s),热请求 0 RTT。"""
|
||||
ent = _local_versions.get(name)
|
||||
now = time.monotonic()
|
||||
if ent is not None and ent[1] > now:
|
||||
return ent[0]
|
||||
c = _client()
|
||||
if c is None:
|
||||
return 0
|
||||
try:
|
||||
v = await c.get(f"ver:{name}")
|
||||
return int(v) if v is not None else 0
|
||||
val = int(v) if v is not None else 0
|
||||
_local_versions[name] = (val, now + _LOCAL_VERSION_TTL)
|
||||
return val
|
||||
except Exception: # noqa: BLE001
|
||||
_bail()
|
||||
return 0
|
||||
@@ -89,7 +141,9 @@ async def bump_version(name: str) -> None:
|
||||
if c is None:
|
||||
return
|
||||
try:
|
||||
await c.incr(f"ver:{name}")
|
||||
v = await c.incr(f"ver:{name}")
|
||||
# 同进程 bump 立即可见(夜同步跑在本进程时响应缓存零陈旧窗口)
|
||||
_local_versions[name] = (int(v), time.monotonic() + _LOCAL_VERSION_TTL)
|
||||
except Exception: # noqa: BLE001
|
||||
_bail()
|
||||
|
||||
|
||||
@@ -10,9 +10,12 @@ class Base(DeclarativeBase):
|
||||
|
||||
|
||||
# echo=False;远程 PG 的空闲连接可能被中间层断开,pre_ping + recycle 自动剔除死连接
|
||||
# pool_size/max_overflow:preview 冷路径每请求开 2 个会话(2 波并发),Redis 故障时全请求变冷,
|
||||
# 默认 5+10 不够并发冷请求用(连接池满了只会排队,不会崩)
|
||||
engine = create_async_engine(
|
||||
settings.database_url, echo=False, future=True,
|
||||
pool_pre_ping=True, pool_recycle=1800,
|
||||
pool_size=10, max_overflow=20,
|
||||
)
|
||||
async_session = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession)
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ Candle 表设计与 TimescaleDB hypertable 完全兼容:将来在目标 PG 库
|
||||
"""
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Boolean, Date, DateTime, Float, ForeignKey, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy import BigInteger, Boolean, Date, DateTime, Float, ForeignKey, Index, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from .db import Base
|
||||
@@ -136,11 +136,14 @@ class AdjFactor(Base):
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
trade_date: Mapped[datetime] = mapped_column(DateTime, index=True)
|
||||
ts_code: Mapped[str] = mapped_column(String(12), index=True)
|
||||
ts_code: Mapped[str] = mapped_column(String(12))
|
||||
adj_factor: Mapped[float] = mapped_column(Float)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("ts_code", "trade_date", name="uq_adj_code_date"),
|
||||
# 覆盖索引:因子查询只取 (trade_date, adj_factor) 两列时走 Index Only Scan,
|
||||
# 免堆访问(adj_factor 堆碎片化严重,见 alembic/versions/20260902_01)
|
||||
Index("ix_adj_code_date_cover", "ts_code", "trade_date", postgresql_include=["adj_factor"]),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ from datetime import datetime, timedelta, timezone
|
||||
import httpx
|
||||
from sqlalchemy import delete, select
|
||||
|
||||
from app.auth import hash_password, token_digest
|
||||
from app.auth import drop_session_cache, hash_password, token_digest
|
||||
from app.db import async_session, engine
|
||||
from app.main import app
|
||||
from app.models import AuthSession, Candle, User
|
||||
@@ -87,6 +87,9 @@ async def main() -> None:
|
||||
).scalar_one()
|
||||
auth_session.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1)
|
||||
await db.commit()
|
||||
# 鉴权会话有 60s 进程内缓存(auth.py),会掩盖库里的过期态;
|
||||
# 清缓存模拟 TTL 已过,验证 require_user 对过期会话本身的判定
|
||||
drop_session_cache(digest=token_digest(token))
|
||||
assert (await client.get("/api/auth/me")).status_code == 401
|
||||
print("expired session: 401")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user