Files
ggx/src/hdiv/core/config.py
T
simon fce725e13c 初始提交:高股息策略研究与回测系统
从 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。
2026-10-03 13:54:56 +08:00

839 lines
29 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""配置加载与严格校验。
设计原则(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}