本轮会话的三项正确性改造(均为「不报错、只让结果静默错」的类型):
1) 修复 stock_daily 量价单位前后不一致
- 现象:2015-2019 存 Tushare 原始单位(手/千元),2020 起存(股/元),2019 同日混合;
而流动性阈值按「元」配置 → 早年门槛实际是「日均成交额 ≥ 200 亿元」,
把 2015-2019 的股票池整体清空(实测 2016/2017/2018 各选出 0 只)。
- 修复:写入端 sync/price.py 统一换算;读取端 units.normalize_ohlcv_units
按行判定并幂等换算(price_history / avg_amount 都走它);
审计新增 UNIT-OHLCV 防回归。
- 效果:2016/2017/2018 的股票池变为 7/11/13 只。
2) 未来函数守卫(单次回测)
- 股票池自带 asof:若晚于回测起点即**拒绝执行**(原先静默冻结套用),
与 walk-forward 已有的拒绝理由一致;确需复现加 --allow-lookahead-universe,
偏差写入 unimplemented_json。
3) 新增实时(PIT)个股画像闸门
- profile/pit.py:每个决策日按当时可见数据重算过去 5 年画像,
惰性(仅买入条件已触发的标的)、面板按 asof 缓存、
规则不含财务指标时不查财报表;被剔除时产出 REJECT + 逐规则留痕。
- 指标定义复用 ProfileBuilder._profile_one(与批量画像逐值等价的回归测试)。
- profile/coverage.py:窗口覆盖率(按交易日历的真实开市天数),
策略新增 entry.profile_gate.min_window_coverage(默认 0,不改变既有行为)。
- core/metrics.py:闸门可用指标的唯一定义(配置期即校验,避免写错指标名静默失效)。
4) 行情回补到 2005(使 5/8/10 年窗口真正完整)
- stock_daily / adjust_factor / daily_basic 补到 2005-01-04;
hd_suspend / hd_limit 补到 2010-01-04。
- 5 年窗口覆盖率:2018-05-18 由 67.0% → 99.1%,2016-12-30 由 39.8% → 99.0%;
残差经逐日与 hd_suspend 交叉核实为真实停牌(16/16 命中)。
- 审计 G2/G3 与断点续传原先用固定阈值(2000 / 1500 只),
会把 2005-2009 的正常数据误判为异常 —— 改为按「当年应有上市股票数」成比例判定。
- 节流修正:daily/adj_factor/daily_basic 限频 480 → 170(实测该 token 约 196/min 即被拒)。
5) 自我声明如实化
- 原先「约束未生效」由「过滤后集合为空」判定,会把「这批股票恰好没停牌」
误报成「hd_suspend 无数据」;改为按表级判定。
- 补齐此前静默的「配置承诺但未实现」项:suspended_rule/limit_up_down_rule 的 defer、
cash_mode=reinvest/reinvest_rule、handle_rights_issue、signal_to_execution、
max_volume_pct、liquidity_limit_pct_adv —— 全部写入 unimplemented_json。
6) 手册:新增 §0「全流程操作(选股 → 画像 → 回测)」置于最前
- 逐步说明「命令做了什么、数据从哪来、落了哪些库、有哪些坑」;
含实时画像闸门 9 问 9 答、未来函数守卫表、成交与成本口径、验证 SQL。
- 修正旧 §2.4 漏传 --universe-run(选了池子却没用于回测);
修正两处声称「停牌顺延」「分红再投资」已实现的相反表述。
测试:403 项全部通过(含新增 test_units.py、test_profile_pit.py、
未实现声明诚实性测试、行序无关性回归测试)。
注意:本提交中 docs/*、README.md、src/hdiv/web/service.py 除本轮修改外,
也含此前遗留的未提交改动(无法按文件切分)。
940 lines
34 KiB
Python
940 lines
34 KiB
Python
"""配置加载与严格校验。
|
||
|
||
设计原则(development-plan.md §10 纪律 2):
|
||
- 任何业务阈值只允许出现在 config/*.yml;
|
||
- 加载时用 pydantic 严格校验,字段名写错 / 类型不符 **直接报错**,不静默取默认值;
|
||
- 每次加载计算 config_hash,写入 run 记录以支持可复现性(plan.md §47 原则 6)。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import hashlib
|
||
import json
|
||
from datetime import date
|
||
from functools import lru_cache
|
||
from pathlib import Path
|
||
from typing import Any, Literal
|
||
|
||
import yaml
|
||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||
|
||
from hdiv.core.errors import ConfigError, ConfigNotFound, SchemaValidationError
|
||
from hdiv.core.paths import config_dir, project_root
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 基类:禁止未声明字段(防拼写错误),禁止类型强转
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class StrictModel(BaseModel):
|
||
"""所有配置模型的基类:多余字段一律报错。"""
|
||
|
||
model_config = ConfigDict(extra="forbid")
|
||
|
||
|
||
DateOrLatest = date | Literal["latest"]
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# datasource.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class DatabaseConfig(StrictModel):
|
||
host: str
|
||
port: int = 3306
|
||
db: str
|
||
user: str
|
||
password_env: str = "MYSQL_PASSWORD"
|
||
charset: str = "utf8mb4"
|
||
pool_size: int = 8
|
||
pool_recycle_sec: int = 3600
|
||
|
||
read_only_tables: list[str] = Field(default_factory=list)
|
||
allow_write_tables: list[str] = Field(default_factory=list)
|
||
backfill_tables: list[str] = Field(default_factory=list)
|
||
allow_backfill: bool = False
|
||
own_prefix: str = "hd_"
|
||
forbidden_tables: list[str] = Field(default_factory=lambda: ["alembic_version"])
|
||
forbid_delete: bool = True
|
||
|
||
@model_validator(mode="after")
|
||
def _check_no_overlap(self) -> DatabaseConfig:
|
||
ro, wr = set(self.read_only_tables), set(self.allow_write_tables)
|
||
both = ro & wr
|
||
if both:
|
||
raise SchemaValidationError(
|
||
f"read_only_tables 与 allow_write_tables 冲突:{sorted(both)}"
|
||
)
|
||
if not self.own_prefix:
|
||
raise SchemaValidationError("own_prefix 不能为空")
|
||
return self
|
||
|
||
|
||
class TushareConfig(StrictModel):
|
||
api_url: str = "http://api.tushare.pro"
|
||
token_env: str = "TUSHARE_TOKEN"
|
||
# 限频是**按接口**计算的(实测 dividend 上限 200/min)。
|
||
# rate_limits 优先,未命中时用 rate_limit_default。
|
||
rate_limit_default: int = 380
|
||
rate_limits: dict[str, int] = Field(default_factory=dict)
|
||
retry: int = 4
|
||
retry_backoff_sec: float = 2.0
|
||
timeout_sec: float = 60.0
|
||
max_retry_backoff_sec: float = 30.0
|
||
# 命中限频后的强制冷却秒数(Tushare 为滑动 60 秒窗口)
|
||
rate_limit_cooldown_sec: float = 62.0
|
||
|
||
def limit_for(self, api_name: str) -> int:
|
||
return int(self.rate_limits.get(api_name, self.rate_limit_default))
|
||
|
||
|
||
class PathsConfig(StrictModel):
|
||
output_dir: str = "output"
|
||
templates_dir: str = "templates"
|
||
assets_dir: str = "assets"
|
||
log_dir: str = "logs"
|
||
|
||
|
||
class SyncConfig(StrictModel):
|
||
"""行情同步的「完整性」判定口径(决定断点续传会不会重拉整段历史)。
|
||
|
||
一个交易日被视为**已完整同步**,要求当日股票数 ≥
|
||
``max(min_symbols_floor, min_symbols_ratio × 当年应有上市股票数)``。
|
||
|
||
**为什么不能只用一个绝对阈值**:A 股 2005 年只有约 1,350 只股票,
|
||
2010 年约 1,700 只。若固定要求 1,500 只,2005-2009 的**每一个交易日**
|
||
都会被判成「未完成」,于是断点续传完全失效 —— 回补中断一次就要从
|
||
第一天重新拉,且每次重跑都会把整段早年历史再拉一遍。
|
||
"""
|
||
|
||
min_symbols_floor: int = 200
|
||
min_symbols_ratio: float = 0.6
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> SyncConfig:
|
||
if self.min_symbols_floor < 1:
|
||
raise SchemaValidationError("sync.min_symbols_floor 必须为正")
|
||
if not 0.0 < self.min_symbols_ratio <= 1.0:
|
||
raise SchemaValidationError("sync.min_symbols_ratio 必须落在 (0, 1]")
|
||
return self
|
||
|
||
|
||
class DataSourceConfig(StrictModel):
|
||
version: int = 1
|
||
database: DatabaseConfig
|
||
tushare: TushareConfig = Field(default_factory=TushareConfig)
|
||
paths: PathsConfig = Field(default_factory=PathsConfig)
|
||
sync: SyncConfig = Field(default_factory=SyncConfig)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# universe.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class EvaluationConfig(StrictModel):
|
||
mode: Literal["point_in_time", "latest"] = "point_in_time"
|
||
asof: date | Literal["latest"] = "latest"
|
||
|
||
|
||
class MarketFilterConfig(StrictModel):
|
||
exchanges: list[str] = Field(default_factory=lambda: ["SZSE", "SSE"])
|
||
markets: list[str] = Field(default_factory=lambda: ["主板", "创业板", "科创板"])
|
||
min_listing_years: int = 10
|
||
min_market_cap: float | None = 50_000_000_000.0
|
||
min_float_market_cap: float | None = None
|
||
min_avg_amount_20d: float | None = 20_000_000.0
|
||
require_trading_on_asof: bool = True
|
||
|
||
|
||
class RiskFilterConfig(StrictModel):
|
||
exclude_st: bool = True
|
||
st_lookback_from_history: bool = True
|
||
exclude_delisting: bool = True
|
||
exclude_suspended: bool = True
|
||
exclude_negative_equity: bool = True
|
||
max_debt_to_assets: float | None = 0.80
|
||
max_pledge_ratio: float | None = None
|
||
exclude_major_litigation: bool = False
|
||
exclude_goodwill_anomaly: bool = False
|
||
|
||
|
||
class DividendFilterConfig(StrictModel):
|
||
yield_source: Literal["computed", "dv_ttm"] = "computed"
|
||
min_dividend_yield: float | None = 0.03
|
||
min_continuous_years: int = 5
|
||
continuous_rule: Literal["cash_div_tax_gt_0"] = "cash_div_tax_gt_0"
|
||
window_years: int = 6
|
||
min_dividend_years_in_window: int = 5
|
||
max_payout_ratio: float | None = 1.00
|
||
min_payout_ratio: float | None = None
|
||
min_dps_cagr_5y: float | None = None
|
||
require_positive_fcf: bool = True
|
||
min_fcf_dividend_cover: float | None = 1.0
|
||
# 验证所需数据缺失时的处理:pass=放行(宽松)/ reject=淘汰(严格)
|
||
on_missing_data: Literal["pass", "reject"] = "pass"
|
||
|
||
@model_validator(mode="after")
|
||
def _check_years(self) -> DividendFilterConfig:
|
||
if self.min_continuous_years > self.window_years:
|
||
raise SchemaValidationError(
|
||
f"dividend.min_continuous_years({self.min_continuous_years}) "
|
||
f"不能大于 window_years({self.window_years})"
|
||
)
|
||
return self
|
||
|
||
|
||
class QualityFilterConfig(StrictModel):
|
||
min_roe_5y_avg: float | None = 0.08
|
||
min_roic_5y_avg: float | None = None
|
||
min_gross_margin: float | None = None
|
||
min_net_margin: float | None = None
|
||
min_ocf_to_profit: float | None = 0.60
|
||
max_debt_to_assets: float | None = None
|
||
min_interest_cover: float | None = None
|
||
exempt_industries: list[str] = Field(default_factory=list)
|
||
pit_rule: Literal["announce_date_le_asof"] = "announce_date_le_asof"
|
||
report_basis: Literal["latest_announced", "ttm"] = "latest_announced"
|
||
|
||
|
||
class UniverseOutputConfig(StrictModel):
|
||
min_members: int = 10
|
||
max_members: int | None = 200
|
||
warn_on_empty: bool = True
|
||
|
||
|
||
class IndustryExemptionsConfig(StrictModel):
|
||
"""按行业豁免特定阈值(金融股口径与实业不同)。"""
|
||
|
||
leverage: list[str] = Field(default_factory=lambda: ["银行", "保险", "证券", "信托"])
|
||
free_cash_flow: list[str] = Field(default_factory=lambda: ["银行", "保险", "证券", "信托"])
|
||
|
||
|
||
class UniverseConfig(StrictModel):
|
||
version: int = 1
|
||
name: str
|
||
evaluation: EvaluationConfig = Field(default_factory=EvaluationConfig)
|
||
industry_exemptions: IndustryExemptionsConfig = Field(
|
||
default_factory=IndustryExemptionsConfig
|
||
)
|
||
market: MarketFilterConfig = Field(default_factory=MarketFilterConfig)
|
||
risk: RiskFilterConfig = Field(default_factory=RiskFilterConfig)
|
||
dividend: DividendFilterConfig = Field(default_factory=DividendFilterConfig)
|
||
quality: QualityFilterConfig = Field(default_factory=QualityFilterConfig)
|
||
output: UniverseOutputConfig = Field(default_factory=UniverseOutputConfig)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# profile.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class DividendYieldVolatilityConfig(StrictModel):
|
||
daily_window: int = 250
|
||
monthly_window: int = 60
|
||
quarterly_window: int = 20
|
||
annual_window: int = 10
|
||
min_obs: int = 20
|
||
|
||
|
||
class ScoreAnchorPercentile(StrictModel):
|
||
best_percentile: float
|
||
worst_percentile: float
|
||
|
||
|
||
class ScoreAnchorAbs(StrictModel):
|
||
metric: str
|
||
best_abs: float
|
||
worst_abs: float
|
||
|
||
|
||
class ScoreAnchorQuality(StrictModel):
|
||
metric: str
|
||
best_abs: float
|
||
worst_abs: float
|
||
|
||
|
||
class ScoreAnchorsConfig(StrictModel):
|
||
dividend_yield: ScoreAnchorPercentile
|
||
valuation: ScoreAnchorPercentile
|
||
financial_quality: ScoreAnchorQuality
|
||
balance_sheet: ScoreAnchorQuality
|
||
dividend_quality: ScoreAnchorQuality
|
||
|
||
|
||
class SafetyMarginConfig(StrictModel):
|
||
mode: Literal["separate", "composite"] = "separate"
|
||
weights: dict[str, float]
|
||
score_bounds: tuple[float, float] = (0.0, 100.0)
|
||
score_anchors: ScoreAnchorsConfig | None = None
|
||
|
||
@model_validator(mode="after")
|
||
def _check_weights(self) -> SafetyMarginConfig:
|
||
if self.mode == "composite":
|
||
total = sum(self.weights.values())
|
||
if abs(total - 1.0) > 1e-6:
|
||
raise SchemaValidationError(
|
||
f"safety_margin.mode=composite 时 weights 之和须为 1.0,当前为 {total}"
|
||
)
|
||
if self.score_bounds[0] >= self.score_bounds[1]:
|
||
raise SchemaValidationError("score_bounds 必须满足 [min, max] 且 min < max")
|
||
return self
|
||
|
||
|
||
class MetricsConfig(StrictModel):
|
||
valuation: list[str] = Field(default_factory=list)
|
||
financial: list[str] = Field(default_factory=list)
|
||
dividend: list[str] = Field(default_factory=list)
|
||
risk: list[str] = Field(default_factory=list)
|
||
return_: list[str] = Field(default_factory=list, alias="return")
|
||
|
||
model_config = ConfigDict(extra="forbid", populate_by_name=True)
|
||
|
||
def all_codes(self) -> list[str]:
|
||
seen: list[str] = []
|
||
for group in (self.valuation, self.financial, self.dividend, self.risk, self.return_):
|
||
for code in group:
|
||
if code not in seen:
|
||
seen.append(code)
|
||
return seen
|
||
|
||
@model_validator(mode="after")
|
||
def _check_duplicates(self) -> MetricsConfig:
|
||
flat = self.valuation + self.financial + self.dividend + self.risk + self.return_
|
||
if len(flat) != len(set(flat)):
|
||
dupes = {c for c in flat if flat.count(c) > 1}
|
||
raise SchemaValidationError(f"profile.metrics 存在重复指标:{sorted(dupes)}")
|
||
return self
|
||
|
||
|
||
class SufficiencyConfig(StrictModel):
|
||
min_history_years_dividend: int = 5
|
||
min_history_years_price: int = 3
|
||
min_dividend_records: int = 5
|
||
|
||
|
||
class TtmDividendConfig(StrictModel):
|
||
window_days: int = 365
|
||
grace_days: int = 45
|
||
# 是否消除除权间隔不规整造成的毛刺(重叠虚高 / 断档虚低)
|
||
smooth_spikes: bool = True
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> TtmDividendConfig:
|
||
if self.window_days <= 0:
|
||
raise SchemaValidationError("ttm_dividend.window_days 必须为正")
|
||
if self.grace_days < 0:
|
||
raise SchemaValidationError("ttm_dividend.grace_days 不能为负")
|
||
return self
|
||
|
||
|
||
class ProfileConfig(StrictModel):
|
||
version: int = 1
|
||
name: str = "default_profile"
|
||
windows_years: list[int] = Field(default_factory=lambda: [5, 8, 10])
|
||
min_obs_days: int = 500
|
||
series_max_points: int = 1500
|
||
series_metrics: list[str] = Field(
|
||
default_factory=lambda: ["dv_yield", "pe_ttm", "pb", "close", "drawdown"]
|
||
)
|
||
ttm_dividend: TtmDividendConfig = Field(default_factory=TtmDividendConfig)
|
||
metrics: MetricsConfig = Field(default_factory=MetricsConfig)
|
||
percentiles: list[float] = Field(default_factory=lambda: [10, 25, 50, 75, 90])
|
||
dividend_yield_volatility: DividendYieldVolatilityConfig = Field(
|
||
default_factory=DividendYieldVolatilityConfig
|
||
)
|
||
safety_margin: SafetyMarginConfig
|
||
sufficiency: SufficiencyConfig = Field(default_factory=SufficiencyConfig)
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> ProfileConfig:
|
||
if not self.percentiles:
|
||
raise SchemaValidationError("profile.percentiles 不能为空")
|
||
if self.percentiles != sorted(self.percentiles):
|
||
raise SchemaValidationError("profile.percentiles 必须升序")
|
||
for p in self.percentiles:
|
||
if not 0 <= p <= 100:
|
||
raise SchemaValidationError(f"分位数越界:{p}")
|
||
if sorted(self.windows_years) != self.windows_years or not self.windows_years:
|
||
raise SchemaValidationError("profile.windows_years 必须为非空升序")
|
||
if self.min_obs_days <= 0:
|
||
raise SchemaValidationError("profile.min_obs_days 必须为正")
|
||
return self
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# cost.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class CommissionConfig(StrictModel):
|
||
rate: float = 0.00025
|
||
min: float = 5.0
|
||
side: Literal["both", "buy", "sell"] = "both"
|
||
|
||
|
||
class SideRateConfig(StrictModel):
|
||
rate: float
|
||
side: Literal["both", "buy", "sell"] = "both"
|
||
|
||
|
||
class SlippageConfig(StrictModel):
|
||
mode: Literal["bps", "fixed", "tick"] = "bps"
|
||
value: float = 10.0
|
||
|
||
|
||
class DividendTaxConfig(StrictModel):
|
||
enabled: bool = True
|
||
holding_based: bool = True
|
||
rates: dict[str, float] = Field(default_factory=dict)
|
||
|
||
|
||
class CostConfig(StrictModel):
|
||
version: int = 1
|
||
model: str = "a_share_default"
|
||
commission: CommissionConfig = Field(default_factory=CommissionConfig)
|
||
stamp_duty: SideRateConfig = Field(
|
||
default_factory=lambda: SideRateConfig(rate=0.0005, side="sell")
|
||
)
|
||
transfer_fee: SideRateConfig = Field(
|
||
default_factory=lambda: SideRateConfig(rate=0.00001, side="both")
|
||
)
|
||
slippage: SlippageConfig = Field(default_factory=SlippageConfig)
|
||
dividend_tax: DividendTaxConfig = Field(default_factory=DividendTaxConfig)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# backtest.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class CapitalConfig(StrictModel):
|
||
initial: float = 1_000_000.0
|
||
currency: str = "CNY"
|
||
|
||
@model_validator(mode="after")
|
||
def _positive(self) -> CapitalConfig:
|
||
if self.initial <= 0:
|
||
raise SchemaValidationError("capital.initial 必须为正")
|
||
return self
|
||
|
||
|
||
class PeriodConfig(StrictModel):
|
||
start: date
|
||
end: DateOrLatest = "latest"
|
||
|
||
@model_validator(mode="after")
|
||
def _order(self) -> PeriodConfig:
|
||
if isinstance(self.end, date) and self.end <= self.start:
|
||
raise SchemaValidationError(
|
||
f"period.end({self.end}) 必须晚于 period.start({self.start})"
|
||
)
|
||
return self
|
||
|
||
|
||
class WalkForwardConfig(StrictModel):
|
||
enabled: bool = True
|
||
scheme: Literal["rolling", "expanding"] = "rolling"
|
||
train_years: int = 5
|
||
test_years: int = 1
|
||
step_months: int = 12
|
||
min_train_years: int = 5
|
||
freeze_params_in_test: bool = True
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> WalkForwardConfig:
|
||
if self.train_years < self.min_train_years:
|
||
raise SchemaValidationError(
|
||
f"walk_forward.train_years({self.train_years}) "
|
||
f"不能小于 min_train_years({self.min_train_years})"
|
||
)
|
||
if self.train_years <= 0 or self.test_years <= 0:
|
||
raise SchemaValidationError("walk_forward 的 train_years / test_years 必须为正")
|
||
if self.step_months <= 0:
|
||
raise SchemaValidationError("walk_forward.step_months 必须为正")
|
||
if not self.freeze_params_in_test:
|
||
raise SchemaValidationError(
|
||
"walk_forward.freeze_params_in_test 必须为 true"
|
||
"(plan.md §25:测试阶段禁止重新调整策略参数)"
|
||
)
|
||
return self
|
||
|
||
|
||
class DividendHandlingConfig(StrictModel):
|
||
cash_mode: Literal["reinvest", "hold", "cash_out"] = "reinvest"
|
||
reinvest_rule: Literal["same_stock_next_open", "portfolio_rebalance"] = (
|
||
"same_stock_next_open"
|
||
)
|
||
apply_dividend_tax: bool = True
|
||
handle_stock_dividend: bool = True
|
||
handle_rights_issue: bool = True
|
||
|
||
|
||
class BenchmarkConfig(StrictModel):
|
||
code: str
|
||
name: str
|
||
|
||
|
||
class FillConfig(StrictModel):
|
||
price: Literal["next_open", "next_close", "same_close"] = "next_open"
|
||
partial_fill: bool = False
|
||
max_volume_pct: float = 0.05
|
||
limit_up_down_rule: Literal["skip", "defer"] = "skip"
|
||
suspended_rule: Literal["skip", "defer"] = "defer"
|
||
|
||
|
||
class ReproConfig(StrictModel):
|
||
record_code_version: bool = True
|
||
record_data_version: bool = True
|
||
seed: int = 42
|
||
|
||
|
||
class ScheduleConfig(StrictModel):
|
||
signal_frequency_months: int = 1
|
||
universe_refresh_months: int = 12
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> ScheduleConfig:
|
||
for n, v in (("signal_frequency_months", self.signal_frequency_months),
|
||
("universe_refresh_months", self.universe_refresh_months)):
|
||
if v <= 0:
|
||
raise SchemaValidationError(f"schedule.{n} 必须为正")
|
||
return self
|
||
|
||
|
||
class PercentileReferenceConfig(StrictModel):
|
||
mode: Literal["rolling", "frozen"] = "rolling"
|
||
lookback_years: int = 5
|
||
min_observations: int = 250
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> PercentileReferenceConfig:
|
||
if self.lookback_years <= 0:
|
||
raise SchemaValidationError("percentile_reference.lookback_years 必须为正")
|
||
return self
|
||
|
||
|
||
class BacktestConfig(StrictModel):
|
||
version: int = 1
|
||
capital: CapitalConfig = Field(default_factory=CapitalConfig)
|
||
period: PeriodConfig
|
||
schedule: ScheduleConfig = Field(default_factory=ScheduleConfig)
|
||
percentile_reference: PercentileReferenceConfig = Field(
|
||
default_factory=PercentileReferenceConfig
|
||
)
|
||
walk_forward: WalkForwardConfig = Field(default_factory=WalkForwardConfig)
|
||
dividend: DividendHandlingConfig = Field(default_factory=DividendHandlingConfig)
|
||
benchmark: list[BenchmarkConfig] = Field(default_factory=list)
|
||
risk_free_rate: float = 0.02
|
||
fill: FillConfig = Field(default_factory=FillConfig)
|
||
reproducibility: ReproConfig = Field(default_factory=ReproConfig)
|
||
|
||
@model_validator(mode="after")
|
||
def _check_benchmarks(self) -> BacktestConfig:
|
||
codes = [b.code for b in self.benchmark]
|
||
if len(codes) != len(set(codes)):
|
||
raise SchemaValidationError(f"benchmark 存在重复代码:{codes}")
|
||
return self
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# report.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class ChartsConfig(StrictModel):
|
||
kline_signals: bool = True
|
||
equity_curve: bool = True
|
||
drawdown: bool = True
|
||
yearly_returns: bool = True
|
||
monthly_heatmap: bool = True
|
||
rolling_metrics: bool = True
|
||
sector_exposure: bool = True
|
||
position_weights: bool = True
|
||
yield_percentile: bool = True
|
||
yield_histogram: bool = True
|
||
financial_trend: bool = True
|
||
sensitivity_heatmap: bool = True
|
||
walkforward_compare: bool = True
|
||
|
||
|
||
class NamingConfig(StrictModel):
|
||
index: str = "index.html"
|
||
audit: str = "data_audit_{date}.html"
|
||
universe: str = "universe_{asof}.html"
|
||
profile: str = "profile_{symbol}_{asof}.html"
|
||
backtest: str = "backtest_{run_id}.html"
|
||
walkforward: str = "walkforward_{wf_id}.html"
|
||
sensitivity: str = "sensitivity_{sens_id}.html"
|
||
|
||
|
||
class DecimalsConfig(StrictModel):
|
||
ratio: int = 4
|
||
money: int = 2
|
||
price: int = 2
|
||
|
||
|
||
class LayoutConfig(StrictModel):
|
||
max_width: int = 1440
|
||
table_page_size: int = 100
|
||
kline_default_days: int = 1500
|
||
normalize_nav: bool = True
|
||
decimals: DecimalsConfig = Field(default_factory=DecimalsConfig)
|
||
|
||
|
||
class ReportConfig(StrictModel):
|
||
version: int = 1
|
||
templates_dir: str = "templates"
|
||
output_dir: str = "output"
|
||
assets_dir: str = "assets"
|
||
theme: Literal["light", "dark", "auto"] = "light"
|
||
offline_assets: bool = True
|
||
asset_mode: Literal["shared", "inline"] = "shared"
|
||
include_sql_provenance: bool = True
|
||
charts: ChartsConfig = Field(default_factory=ChartsConfig)
|
||
naming: NamingConfig = Field(default_factory=NamingConfig)
|
||
layout: LayoutConfig = Field(default_factory=LayoutConfig)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# strategy/*.yml
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class StrategyMeta(StrictModel):
|
||
id: str
|
||
name: str
|
||
version: str
|
||
status: Literal[
|
||
"DRAFT", "RESEARCH", "BACKTEST", "VALIDATED", "PAPER", "LIVE", "ARCHIVED"
|
||
] = "DRAFT"
|
||
description: str = ""
|
||
|
||
|
||
class UniverseRef(StrictModel):
|
||
ref: str = "config/universe.yml"
|
||
override: dict[str, Any] = Field(default_factory=dict)
|
||
|
||
|
||
class ScaleStep(StrictModel):
|
||
percentile: float
|
||
weight: float
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> ScaleStep:
|
||
if not 0 <= self.percentile <= 100:
|
||
raise SchemaValidationError(f"scale 步骤的分位数越界:{self.percentile}")
|
||
if not 0.0 <= self.weight <= 1.0:
|
||
raise SchemaValidationError(f"scale 步骤的权重越界:{self.weight}")
|
||
return self
|
||
|
||
|
||
class ProfileGateRule(StrictModel):
|
||
"""一条实时画像闸门规则。
|
||
|
||
语义:``<metric> 的 <stat> <op> <value>`` 必须成立,否则不买。
|
||
指标名与分位可用性在**配置期**校验 —— 写错一个指标名若拖到运行时,
|
||
只会得到「无法验证 → 保守不买」,表现为策略再也不交易,极难定位。
|
||
"""
|
||
|
||
metric: str
|
||
stat: Literal["current_value", "current_percentile"] = "current_value"
|
||
op: Literal[">=", "<=", ">", "<"] = ">="
|
||
value: float
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> ProfileGateRule:
|
||
from hdiv.core.metrics import GATE_METRICS, PERCENTILE_METRICS
|
||
|
||
if self.metric not in GATE_METRICS:
|
||
head = self.metric.split("_")[0]
|
||
near = sorted(m for m in GATE_METRICS if head and head in m)
|
||
raise SchemaValidationError(
|
||
f"profile_gate 规则引用了未知指标 {self.metric!r}。"
|
||
f"可选指标见 hdiv/core/metrics.py;相近的有 {near[:6]}"
|
||
)
|
||
if self.stat == "current_percentile" and self.metric not in PERCENTILE_METRICS:
|
||
raise SchemaValidationError(
|
||
f"{self.metric} 是标量指标,没有历史分位,不能用 "
|
||
f"stat=current_percentile。有分位的指标:{sorted(PERCENTILE_METRICS)}"
|
||
)
|
||
return self
|
||
|
||
|
||
class ProfileGateConfig(StrictModel):
|
||
"""实时(PIT)个股画像闸门。
|
||
|
||
在每个决策日、**买入条件已经触发之后**,用「当时可见的数据」重算画像,
|
||
不通过的票直接剔除。跨股票共享的面板按时点缓存,代价与「触发次数」成正比,
|
||
而不是与「回测区间 × 股票数」成正比。
|
||
"""
|
||
|
||
#: 是否启用。关闭时回测行为与启用前完全一致(可用于复现历史结果)
|
||
enabled: bool = False
|
||
#: 画像统计窗口(年)。0 = 全历史;其余必须是 config/profile.yml 的 windows_years 之一
|
||
window_years: int = 5
|
||
#: 数据缺失/样本不足(无法验证)时:reject = 保守不买,pass = 放行
|
||
on_unverifiable: Literal["reject", "pass"] = "reject"
|
||
#: 窗口**实际覆盖率**下限(1.0 = 名义 5 年就必须真有 5 年数据)。
|
||
#: 0 = 不因覆盖率淘汰(默认,保持改造前行为)。
|
||
#:
|
||
#: 为什么需要它:`window_slice(asof, 5)` 只是「把已有数据切成最近 5 年」,
|
||
#: 数据起点晚于窗口左端时窗口会被静默截短 —— 实测 2018-05-18 的「5 年」
|
||
#: 窗口只有 3.4 年(817/1215 个交易日,67%),而画像仍报 OK。
|
||
min_window_coverage: float = 0.0
|
||
rules: list[ProfileGateRule] = Field(default_factory=list)
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> ProfileGateConfig:
|
||
if self.window_years < 0:
|
||
raise SchemaValidationError("profile_gate.window_years 不能为负")
|
||
if not 0.0 <= self.min_window_coverage <= 1.0:
|
||
raise SchemaValidationError(
|
||
f"profile_gate.min_window_coverage 必须落在 [0, 1],"
|
||
f"当前 {self.min_window_coverage}"
|
||
)
|
||
if self.enabled and not self.rules:
|
||
raise SchemaValidationError(
|
||
"profile_gate.enabled=true 但 rules 为空 —— 空闸门等于每次都要"
|
||
"算一遍画像再无条件放行。请补齐规则,或把 enabled 设为 false。"
|
||
)
|
||
return self
|
||
|
||
|
||
class EntryConfig(StrictModel):
|
||
yield_percentile: float
|
||
require_universe_pass: bool = True
|
||
require_risk_pass: bool = True
|
||
scale_in: list[ScaleStep] = Field(default_factory=list)
|
||
profile_gate: ProfileGateConfig = Field(default_factory=ProfileGateConfig)
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> EntryConfig:
|
||
if not 0 <= self.yield_percentile <= 100:
|
||
raise SchemaValidationError(f"entry.yield_percentile 越界:{self.yield_percentile}")
|
||
if self.scale_in:
|
||
for a, b in zip(self.scale_in, self.scale_in[1:], strict=False):
|
||
if b.percentile <= a.percentile:
|
||
raise SchemaValidationError("entry.scale_in 必须按 percentile 严格升序")
|
||
if b.weight < a.weight:
|
||
raise SchemaValidationError("entry.scale_in 的 weight 必须随分位单调不减")
|
||
if self.scale_in[0].percentile != self.yield_percentile:
|
||
raise SchemaValidationError(
|
||
f"entry.scale_in 首个分位({self.scale_in[0].percentile}) "
|
||
f"必须等于 entry.yield_percentile({self.yield_percentile})"
|
||
)
|
||
return self
|
||
|
||
|
||
class ExitConfig(StrictModel):
|
||
yield_percentile: float
|
||
scale_out: list[ScaleStep] = Field(default_factory=list)
|
||
stop_loss_pct: float | None = None
|
||
max_holding_days: int | None = None
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> ExitConfig:
|
||
if not 0 <= self.yield_percentile <= 100:
|
||
raise SchemaValidationError(f"exit.yield_percentile 越界:{self.yield_percentile}")
|
||
if self.scale_out:
|
||
for a, b in zip(self.scale_out, self.scale_out[1:], strict=False):
|
||
if b.percentile >= a.percentile:
|
||
raise SchemaValidationError("exit.scale_out 必须按 percentile 严格降序")
|
||
if b.weight > a.weight:
|
||
raise SchemaValidationError("exit.scale_out 的 weight 必须随分位单调不减")
|
||
if self.scale_out[-1].percentile != self.yield_percentile:
|
||
raise SchemaValidationError(
|
||
f"exit.scale_out 末个分位({self.scale_out[-1].percentile}) "
|
||
f"必须等于 exit.yield_percentile({self.yield_percentile})"
|
||
)
|
||
if abs(self.scale_out[-1].weight) > 1e-9:
|
||
raise SchemaValidationError("exit.scale_out 末步 weight 必须为 0(清仓)")
|
||
return self
|
||
|
||
|
||
class PositionConfig(StrictModel):
|
||
max_position: float = 0.10
|
||
min_position: float = 0.01
|
||
sector_max_position: float = 0.25
|
||
max_holdings: int = 20
|
||
weight_scheme: Literal["equal", "score", "inverse_vol"] = "equal"
|
||
|
||
@model_validator(mode="after")
|
||
def _check(self) -> PositionConfig:
|
||
for name, v in (
|
||
("max_position", self.max_position),
|
||
("min_position", self.min_position),
|
||
("sector_max_position", self.sector_max_position),
|
||
):
|
||
if not 0.0 < v <= 1.0:
|
||
raise SchemaValidationError(f"position.{name} 必须落在 (0, 1],当前 {v}")
|
||
if self.min_position > self.max_position:
|
||
raise SchemaValidationError("position.min_position 不能大于 max_position")
|
||
if self.max_holdings <= 0:
|
||
raise SchemaValidationError("position.max_holdings 必须为正")
|
||
if self.max_position * self.max_holdings < 1.0 - 1e-9:
|
||
# 不阻断,但这是常见配置错误:单股上限 × 持仓数 < 1 会导致仓位无法打满
|
||
pass
|
||
return self
|
||
|
||
|
||
class RiskConfig(StrictModel):
|
||
max_portfolio_drawdown: float | None = 0.20
|
||
max_single_drawdown: float | None = 0.25
|
||
liquidity_limit_pct_adv: float | None = 0.05
|
||
|
||
|
||
class ExecutionConfig(StrictModel):
|
||
signal_to_execution: Literal["next_open", "next_close", "same_close"] = "next_open"
|
||
limit_up_down_rule: Literal["skip", "defer"] = "skip"
|
||
suspended_rule: Literal["skip", "defer"] = "defer"
|
||
|
||
|
||
class CostRef(StrictModel):
|
||
ref: str = "config/cost.yml"
|
||
|
||
|
||
class StrategyConfig(StrictModel):
|
||
strategy: StrategyMeta
|
||
universe: UniverseRef = Field(default_factory=UniverseRef)
|
||
entry: EntryConfig
|
||
exit: ExitConfig
|
||
position: PositionConfig = Field(default_factory=PositionConfig)
|
||
risk: RiskConfig = Field(default_factory=RiskConfig)
|
||
execution: ExecutionConfig = Field(default_factory=ExecutionConfig)
|
||
cost: CostRef = Field(default_factory=CostRef)
|
||
|
||
@model_validator(mode="after")
|
||
def _check_entry_exit(self) -> StrategyConfig:
|
||
if self.entry.yield_percentile <= self.exit.yield_percentile:
|
||
raise SchemaValidationError(
|
||
f"entry.yield_percentile({self.entry.yield_percentile}) 必须大于 "
|
||
f"exit.yield_percentile({self.exit.yield_percentile})"
|
||
)
|
||
return self
|
||
|
||
def param_map(self) -> dict[str, Any]:
|
||
"""扁平化参数路径 → 值,写入 hd_strategy_param,便于跨版本 diff 与敏感性扫描。"""
|
||
out: dict[str, Any] = {}
|
||
|
||
def walk(prefix: str, obj: Any) -> None:
|
||
if isinstance(obj, BaseModel):
|
||
for name in type(obj).model_fields:
|
||
walk(f"{prefix}.{name}" if prefix else name, getattr(obj, name))
|
||
elif isinstance(obj, list):
|
||
for i, item in enumerate(obj):
|
||
walk(f"{prefix}[{i}]", item)
|
||
elif isinstance(obj, dict):
|
||
for k, v in obj.items():
|
||
walk(f"{prefix}.{k}" if prefix else str(k), v)
|
||
else:
|
||
out[prefix] = obj
|
||
|
||
walk("", self)
|
||
return out
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 加载与哈希
|
||
# ---------------------------------------------------------------------------
|
||
|
||
_MODEL_BY_NAME: dict[str, type[StrictModel]] = {
|
||
"datasource": DataSourceConfig,
|
||
"universe": UniverseConfig,
|
||
"profile": ProfileConfig,
|
||
"cost": CostConfig,
|
||
"backtest": BacktestConfig,
|
||
"report": ReportConfig,
|
||
}
|
||
|
||
|
||
def resolve_strategy_path(rel: str | Path) -> Path:
|
||
"""策略路径解析,按以下顺序尝试:
|
||
|
||
1. 绝对路径 → 原样
|
||
2. 项目根相对(如 ``config/strategy/high_dividend_v1.yml``)
|
||
3. ``config/`` 相对(如 ``strategy/high_dividend_v1.yml``)
|
||
4. ``config/strategy/`` 下的裸文件名(如 ``high_dividend_v1.yml``)
|
||
"""
|
||
p = Path(rel)
|
||
if p.is_absolute():
|
||
return p
|
||
candidates = [
|
||
project_root() / p,
|
||
config_dir() / p,
|
||
config_dir() / "strategy" / p,
|
||
]
|
||
for cand in candidates:
|
||
if cand.is_file():
|
||
return cand
|
||
raise ConfigNotFound(
|
||
"策略配置不存在,已尝试:\n " + "\n ".join(str(c) for c in candidates)
|
||
)
|
||
|
||
|
||
def _read_yaml(path: Path) -> dict[str, Any]:
|
||
if not path.is_file():
|
||
raise ConfigNotFound(f"配置文件不存在:{path}")
|
||
try:
|
||
raw = yaml.safe_load(path.read_text(encoding="utf-8"))
|
||
except yaml.YAMLError as exc: # pragma: no cover - 由 YAML 库决定
|
||
raise ConfigError(f"YAML 解析失败:{path}\n{exc}") from exc
|
||
if raw is None:
|
||
raise ConfigError(f"配置文件为空:{path}")
|
||
if not isinstance(raw, dict):
|
||
raise ConfigError(f"配置文件顶层必须是映射(dict):{path}")
|
||
return raw
|
||
|
||
|
||
def load_config(name: str, path: str | Path | None = None) -> Any:
|
||
"""按配置名加载并校验。
|
||
|
||
name ∈ {datasource, universe, profile, cost, backtest, report}
|
||
或 'strategy:<相对路径>' 形式加载策略文件。
|
||
"""
|
||
if name.startswith("strategy:"):
|
||
rel = name.split(":", 1)[1]
|
||
p = resolve_strategy_path(rel)
|
||
raw = _read_yaml(p)
|
||
try:
|
||
return StrategyConfig.model_validate(raw)
|
||
except Exception as exc:
|
||
raise ConfigError(f"策略配置校验失败:{p}\n{exc}") from exc
|
||
|
||
if name not in _MODEL_BY_NAME:
|
||
raise ConfigError(f"未知配置名:{name}(可用:{sorted(_MODEL_BY_NAME)})")
|
||
p = Path(path) if path else config_dir() / f"{name}.yml"
|
||
raw = _read_yaml(p)
|
||
try:
|
||
return _MODEL_BY_NAME[name].model_validate(raw)
|
||
except Exception as exc:
|
||
raise ConfigError(f"配置校验失败:{p}\n{exc}") from exc
|
||
|
||
|
||
@lru_cache(maxsize=32)
|
||
def _cached_config(name: str) -> Any:
|
||
return load_config(name)
|
||
|
||
|
||
def get_config(name: str) -> Any:
|
||
"""带缓存的配置读取(同进程内只解析一次)。"""
|
||
return _cached_config(name)
|
||
|
||
|
||
def config_hash(cfg: BaseModel | Any) -> str:
|
||
"""配置指纹:用于 run 记录与可复现性校验。"""
|
||
if isinstance(cfg, BaseModel):
|
||
payload = cfg.model_dump(mode="json", exclude_none=False)
|
||
elif isinstance(cfg, dict):
|
||
payload = cfg
|
||
else:
|
||
payload = {"value": str(cfg)}
|
||
blob = json.dumps(payload, sort_keys=True, ensure_ascii=False, default=str)
|
||
return hashlib.sha256(blob.encode("utf-8")).hexdigest()[:32]
|
||
|
||
|
||
def load_all(names: tuple[str, ...] = ("datasource", "universe", "profile", "cost", "backtest", "report")) -> dict[str, Any]:
|
||
"""一次性加载全部主配置。"""
|
||
return {n: load_config(n) for n in names}
|