Skip to content

Commit 8aef435

Browse files
authored
Merge pull request #156 from shy3130/feat/regime-filter
feat(regime): 市场状态识别系统 + 叠加策略回测过滤器
2 parents 312c02f + 4ee55e4 commit 8aef435

14 files changed

Lines changed: 1341 additions & 2 deletions

File tree

backend/app/api/backtest.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -207,6 +207,7 @@ class StrategyBacktestRequest(BaseModel):
207207
holding_days: int = 5
208208
asset_type: str = "stock"
209209
minute_fill: bool = False
210+
regime_filter: dict | None = None
210211

211212

212213
@router.post("/strategy/run")
@@ -241,6 +242,7 @@ def strategy_run(req: StrategyBacktestRequest, request: Request):
241242
holding_days=req.holding_days,
242243
asset_type=req.asset_type,
243244
minute_fill=req.minute_fill,
245+
regime_filter=req.regime_filter,
244246
)
245247
task = make_worker_task("backtest", settings.data_dir, cfg)
246248
return run_worker_task(task)
@@ -315,8 +317,9 @@ def _make_job_key(
315317
commission_pct: float | None = None, stamp_tax_pct: float | None = None,
316318
asset_type: str = "stock",
317319
minute_fill: bool = False,
320+
regime_filter: str | None = None,
318321
) -> str:
319-
raw = f"{strategy_id}|{symbols}|{start}|{end}|{matching}|{entry_fill}|{exit_fill}|{fees_pct}|{slippage_bps}|{max_positions}|{max_exposure_pct}|{initial_capital}|{position_sizing}|{params}|{overrides}|{mode}|{holding_days}|{commission_pct}|{stamp_tax_pct}|{asset_type}|{minute_fill}"
322+
raw = f"{strategy_id}|{symbols}|{start}|{end}|{matching}|{entry_fill}|{exit_fill}|{fees_pct}|{slippage_bps}|{max_positions}|{max_exposure_pct}|{initial_capital}|{position_sizing}|{params}|{overrides}|{mode}|{holding_days}|{commission_pct}|{stamp_tax_pct}|{asset_type}|{minute_fill}|{regime_filter}"
320323
return hashlib.md5(raw.encode()).hexdigest()[:12]
321324

322325

@@ -344,6 +347,7 @@ async def strategy_stream(
344347
holding_days: int = 5,
345348
asset_type: str = "stock",
346349
minute_fill: bool = False,
350+
regime_filter: str | None = None,
347351
):
348352
"""SSE 流式策略回测: 实时推送进度, 完成后推送结果, 支持重连 (刷新/切页后恢复)。
349353
@@ -383,6 +387,7 @@ async def strategy_stream(
383387
commission_pct, stamp_tax_pct,
384388
asset_type=asset_type,
385389
minute_fill=minute_fill,
390+
regime_filter=regime_filter,
386391
)
387392

388393
_cleanup_stale_jobs()
@@ -443,6 +448,7 @@ async def event_generator():
443448
holding_days=int(holding_days),
444449
asset_type=asset_type,
445450
minute_fill=minute_fill,
451+
regime_filter=json.loads(regime_filter) if regime_filter else None,
446452
)
447453

448454
def _run_backtest():

backend/app/api/regime.py

Lines changed: 147 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,147 @@
1+
"""市场环境(regime) API — 时序查询 + 手动重算。
2+
3+
装配逻辑在 app.services.regime_builder(纯函数), API 层薄壳 + TTL 缓存。
4+
"""
5+
from __future__ import annotations
6+
7+
import threading
8+
import time
9+
from datetime import date
10+
from typing import Any
11+
12+
from fastapi import APIRouter, Query, Request
13+
14+
from app.services import regime_builder
15+
16+
router = APIRouter(prefix="/api/regime", tags=["regime"])
17+
18+
_CACHE_TTL = 5.0
19+
_cache: dict[str, Any] | None = None
20+
_cache_ts: float = 0.0
21+
_cache_lock = threading.Lock()
22+
23+
24+
def invalidate_regime_cache() -> None:
25+
"""清空 regime 查询缓存。批算/重算后调用。"""
26+
global _cache, _cache_ts
27+
with _cache_lock:
28+
_cache = None
29+
_cache_ts = 0.0
30+
31+
32+
def _data_dir(request: Request) -> Any:
33+
return request.app.state.repo.store.data_dir
34+
35+
36+
def _df_to_records(df) -> list[dict]:
37+
"""polars DataFrame → JSON 安全的 list[dict](date 转 ISO 字符串)。"""
38+
if df is None or df.is_empty():
39+
return []
40+
records = []
41+
for r in df.to_dicts():
42+
if "date" in r and r["date"] is not None:
43+
r["date"] = str(r["date"])
44+
records.append(r)
45+
return records
46+
47+
48+
@router.get("/history")
49+
def regime_history(
50+
request: Request,
51+
start: date | None = Query(None),
52+
end: date | None = Query(None),
53+
limit: int = Query(120, ge=1, le=1000),
54+
):
55+
"""历史环境时序(含状态/指标)。默认最近 N 天。"""
56+
global _cache, _cache_ts
57+
cache_key = f"hist|{start}|{end}|{limit}"
58+
with _cache_lock:
59+
if (
60+
_cache is not None
61+
and _cache.get("key") == cache_key
62+
and (time.time() - _cache_ts) < _CACHE_TTL
63+
):
64+
return _cache["data"]
65+
66+
df = regime_builder.load_regime_history(_data_dir(request))
67+
if df.is_empty():
68+
result: dict = {"rows": [], "total": 0}
69+
else:
70+
if start:
71+
df = df.filter(pl_col_date(df, ">=", start))
72+
if end:
73+
df = df.filter(pl_col_date(df, "<=", end))
74+
df = df.sort("date", descending=True).head(limit).sort("date")
75+
rows = _df_to_records(df)
76+
result = {"rows": rows, "total": len(rows)}
77+
78+
with _cache_lock:
79+
_cache = {"key": cache_key, "data": result}
80+
_cache_ts = time.time()
81+
return result
82+
83+
84+
def pl_col_date(df, op: str, value: date):
85+
"""polars 日期过滤辅助(避免重复 import)。"""
86+
import polars as pl
87+
88+
col = pl.col("date")
89+
return col >= value if op == ">=" else col <= value
90+
91+
92+
@router.get("/latest")
93+
def regime_latest(request: Request):
94+
"""最新一日环境(轻量)。"""
95+
df = regime_builder.load_regime_history(_data_dir(request))
96+
if df.is_empty():
97+
return {"row": None}
98+
latest = df.sort("date", descending=True).head(1)
99+
rows = _df_to_records(latest)
100+
return {"row": rows[0] if rows else None}
101+
102+
103+
@router.get("/states")
104+
def regime_states(
105+
request: Request,
106+
days: int = Query(60, ge=1, le=1000),
107+
):
108+
"""状态分布统计(各状态天数/占比)。"""
109+
df = regime_builder.load_regime_history(_data_dir(request))
110+
if df.is_empty():
111+
return {"distribution": [], "days": 0}
112+
df = df.sort("date", descending=True).head(days)
113+
total = df.height
114+
counts = df.group_by("state").len().sort("len", descending=True)
115+
distribution = [
116+
{
117+
"state": r["state"],
118+
"label": regime_builder.STATE_LABELS.get(r["state"], r["state"]),
119+
"count": r["len"],
120+
"pct": round(r["len"] / total * 100, 1) if total else 0,
121+
}
122+
for r in counts.to_dicts()
123+
]
124+
return {"distribution": distribution, "days": total}
125+
126+
127+
@router.get("/coverage")
128+
def regime_coverage(request: Request):
129+
"""regime 数据覆盖元信息(供数据画像)。"""
130+
return regime_builder.get_regime_coverage(_data_dir(request))
131+
132+
133+
@router.post("/recompute")
134+
def regime_recompute(request: Request, start: date | None = None, end: date | None = None):
135+
"""手动触发重算(全量或指定区间)。管理员操作。"""
136+
repo = request.app.state.repo
137+
data_dir = _data_dir(request)
138+
end = end or date.today()
139+
if start is None:
140+
# 全量: 从 enriched 最早日算到今天
141+
new_rows = regime_builder.compute_regime_incremental(repo, data_dir, today=end)
142+
else:
143+
new_rows = regime_builder.run_regime_batch(repo, start=start, end=end)
144+
if not new_rows.is_empty():
145+
regime_builder.upsert_regime_history(data_dir, new_rows)
146+
invalidate_regime_cache()
147+
return {"ok": True, "computed": new_rows.height if not new_rows.is_empty() else 0}

backend/app/backtest/strategy.py

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
from collections.abc import Callable, Mapping
1414
from dataclasses import dataclass, field
1515
from datetime import date, timedelta
16+
from pathlib import Path
1617
from typing import Literal
1718

1819
import numpy as np
@@ -473,6 +474,9 @@ class StrategyBacktestConfig:
473474
holding_days: int = 5
474475
# 分钟K精确成交: 开启后用当日分钟K确定穿越价/VWAP (需 Pro+ 分钟K能力)
475476
minute_fill: bool = False
477+
# 市场环境过滤: {"states": ["strong",...], "min_score": 60}。
478+
# 强制 T-1: regime[T-1] 决定 entry[T](防未来函数)。None=不过滤。
479+
regime_filter: dict | None = None
476480

477481
def __post_init__(self) -> None:
478482
if self.entry_fill is None:
@@ -846,6 +850,13 @@ def prepare_matrix_optimization(
846850
first.start,
847851
first.end,
848852
)
853+
# 市场环境过滤(优化器共享, 用首个 config 的 regime_filter)
854+
_rm = self._build_regime_mask(
855+
market_data.timestamp_labels, first.regime_filter,
856+
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
857+
)
858+
if _rm is not None:
859+
entry_time_mask = entry_time_mask & _rm
849860
exit_time_mask = self._matrix_date_range_mask(
850861
market_data.timestamp_labels,
851862
first.start,
@@ -1127,6 +1138,13 @@ def _err(msg: str) -> StrategyBacktestResult:
11271138
config.start,
11281139
config.end,
11291140
)
1141+
# 市场环境过滤(强制 T-1): 只叠加 entry, 不影响 exit
1142+
_rm = self._build_regime_mask(
1143+
market_data.timestamp_labels, config.regime_filter,
1144+
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
1145+
)
1146+
if _rm is not None:
1147+
entry_time_mask = entry_time_mask & _rm
11301148
exit_time_mask = self._matrix_date_range_mask(
11311149
market_data.timestamp_labels,
11321150
config.start,
@@ -1220,6 +1238,12 @@ def _err(msg: str) -> StrategyBacktestResult:
12201238
config.start,
12211239
config.end,
12221240
)
1241+
_rm = self._build_regime_mask(
1242+
market_data.timestamp_labels, config.regime_filter,
1243+
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
1244+
)
1245+
if _rm is not None:
1246+
entry_time_mask = entry_time_mask & _rm
12231247
exit_time_mask = self._matrix_date_range_mask(
12241248
market_data.timestamp_labels,
12251249
config.start,
@@ -1620,6 +1644,56 @@ def _matrix_date_range_mask(
16201644
count=len(timestamp_labels),
16211645
)
16221646

1647+
@staticmethod
1648+
def _build_regime_mask(
1649+
timestamp_labels: tuple[str, ...],
1650+
regime_filter: dict | None,
1651+
data_dir: Path | None,
1652+
) -> np.ndarray | None:
1653+
"""构造逐日 regime mask。强制 T-1 防未来函数: regime[T-1] 决定 entry[T]。
1654+
1655+
timestamp_labels[i] 的入场资格 = 它的"前一交易日"的 regime 是否满足条件。
1656+
"前一交易日"用 timestamp_labels 自身的顺序确定(回测时间轴上的前一天)。
1657+
边界: 首日无前一日环境 → 默认允许(不阻断)。
1658+
regime_filter 为 None 或无 regime 数据时返回 None(不过滤)。
1659+
"""
1660+
if not regime_filter or data_dir is None:
1661+
return None
1662+
allowed_states = set(regime_filter.get("states") or [])
1663+
min_score = regime_filter.get("min_score")
1664+
if not allowed_states and min_score is None:
1665+
return None
1666+
1667+
from app.services import regime_builder
1668+
regime_df = regime_builder.load_regime_history(data_dir)
1669+
if regime_df.is_empty():
1670+
return None
1671+
1672+
# 构建 date(ISO) → (state, score) 映射
1673+
regime_map: dict[str, tuple[str, int]] = {}
1674+
for r in regime_df.iter_rows(named=True):
1675+
d = r.get("date")
1676+
ds = str(d)[:10] if d is not None else None
1677+
if ds:
1678+
regime_map[ds] = (str(r.get("state", "")), int(r.get("score", 0) or 0))
1679+
1680+
# 对每个 label, 找它的前一交易日的 regime(timestamp_labels 顺序里的前一天)
1681+
n = len(timestamp_labels)
1682+
mask = np.ones(n, dtype=bool) # 默认允许
1683+
for i in range(1, n):
1684+
prev_label = timestamp_labels[i - 1][:10]
1685+
entry = regime_map.get(prev_label)
1686+
if entry is None:
1687+
continue # 无前一日环境数据 → 允许(不阻断)
1688+
state, score = entry
1689+
ok = True
1690+
if allowed_states and state not in allowed_states:
1691+
ok = False
1692+
if min_score is not None and score < min_score:
1693+
ok = False
1694+
mask[i] = ok
1695+
return mask
1696+
16231697
def _build_candidate_filter_mask(
16241698
self,
16251699
panel: pl.DataFrame,

backend/app/jobs/daily_pipeline.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -516,6 +516,24 @@ def _minute_chunk_progress(cur: int, tot: int, seg_label: str = "") -> None:
516516
else:
517517
logger.info("sync_minute skipped: user disabled")
518518

519+
# Step 2.6: 市场环境(regime) 增量计算 — enriched 已就绪后聚合环境指标。
520+
# 双检测(缺口+stale), 自动补算遗漏/被覆写的日。软失败: 不阻断主管道。
521+
regime_days = 0
522+
try:
523+
emit("compute_regime", 90, "计算市场环境…")
524+
from app.services import regime_builder
525+
from app.api.regime import invalidate_regime_cache
526+
new_regime = regime_builder.compute_regime_incremental(repo, repo.store.data_dir)
527+
regime_days = new_regime.height if not new_regime.is_empty() else 0
528+
if regime_days:
529+
invalidate_regime_cache()
530+
logger.info("compute_regime: %d days", regime_days)
531+
emit("compute_regime", 92, f"市场环境 {regime_days} 天")
532+
except Exception as e: # noqa: BLE001
533+
logger.warning("compute_regime failed (soft): %s", e)
534+
stage_errors.append(f"compute_regime: {e}")
535+
skipped.append("regime")
536+
519537
# Step 3: 刷新视图
520538
emit("refresh_views", 95, "刷新 DuckDB 视图…")
521539
_refresh_views(repo)
@@ -534,6 +552,7 @@ def _minute_chunk_progress(cur: int, tot: int, seg_label: str = "") -> None:
534552
"etf_daily_rows": written_etf_daily,
535553
"etf_adj_factor_symbols": etf_adj_symbols,
536554
"minute_rows": written_minute,
555+
"regime_days": regime_days,
537556
"lagging_symbols": len(lagging_symbols),
538557
"skipped_stages": skipped,
539558
"stage_errors": stage_errors,

backend/app/main.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212
from fastapi.staticfiles import StaticFiles
1313

1414
from app import __version__
15-
from app.api import analysis, auth as auth_api, backtest, data, ext_data, financials, indices, intraday, kline, market_recap, monitor_rules, alerts, overview, pipeline, rps, screener, settings as settings_api, signals, stock_analysis, strategy, watchlist
15+
from app.api import analysis, auth as auth_api, backtest, data, ext_data, financials, indices, intraday, kline, market_recap, monitor_rules, alerts, overview, pipeline, regime, rps, screener, settings as settings_api, signals, stock_analysis, strategy, watchlist
1616
from app.api.routes import router as core_router
1717
from app.config import settings
1818
from app.jobs import daily_pipeline
@@ -338,6 +338,7 @@ async def auth_middleware(request: Request, call_next):
338338
app.include_router(intraday.router)
339339
app.include_router(indices.router)
340340
app.include_router(overview.router)
341+
app.include_router(regime.router)
341342
app.include_router(analysis.router)
342343
app.include_router(pipeline.router)
343344
app.include_router(data.router)

0 commit comments

Comments
 (0)