初始提交:高股息策略研究与回测系统
从 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:
@@ -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}
|
||||
Reference in New Issue
Block a user