"""配置加载与严格校验。 设计原则(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: 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}