初始提交:高股息策略研究与回测系统

从 Point-in-Time 股票筛选到统一 Web 前端的完整链路:
筛选 → 画像 → 策略 → 回测 → Walk-forward → 绩效分析 → 报告/前端。

架构
- 数据层与策略层分离;策略代码不写 SQL,只经 data/repo.py 取数
- 所有业务阈值集中在 config/*.yml,代码零硬编码(字段写错直接报错)
- 报告只做「run_id → SQL → 渲染」,不做任何计算,数字可追溯
- 前后端分离:output/ 静态站点 + hdiv web 提供的 REST API

数据安全
- 只增不删:SQL 钩子拦截 DELETE/DROP/TRUNCATE,并有源码扫描测试守护
- qlib 原有表只读,本项目数据写入 hd_ 前缀表
- 回补使用 INSERT IGNORE,保证既有行零改动
- .env 存密钥且已 gitignore;output/、logs/、.venv/ 不入库

交付物
- 30 张 hd_* 表、7 个 YAML 配置、283 项自动化测试
- 统一 Web 前端(hash 路由 SPA)+ nginx 部署配置与 launchd 托管脚本

如实声明的限制
- 策略缺少稳定的样本外超额收益(Walk-forward 7 窗口均值 -0.95%,
  基准 +2.29%);其价值体现在回撤控制,而非超额收益
- 涨跌停/停牌约束仅覆盖 2019 年起;index_weight 尚未填充
- AI Agent 层(plan.md 第四版 P8)未实现

详见 docs/user-guide.md 与 docs/implementation-status.md。
This commit is contained in:
2026-10-03 13:54:56 +08:00
commit fce725e13c
127 changed files with 28001 additions and 0 deletions
+838
View File
@@ -0,0 +1,838 @@
"""配置加载与严格校验。
设计原则(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 DataSourceConfig(StrictModel):
version: int = 1
database: DatabaseConfig
tushare: TushareConfig = Field(default_factory=TushareConfig)
paths: PathsConfig = Field(default_factory=PathsConfig)
# ---------------------------------------------------------------------------
# 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
@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
@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 EntryConfig(StrictModel):
yield_percentile: float
require_universe_pass: bool = True
require_risk_pass: bool = True
scale_in: list[ScaleStep] = Field(default_factory=list)
@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}