Files
ggx/src/hdiv/core/config.py
T
simon 14ec0c6c86 修复:量价单位 / 未来函数守卫 / 实时画像闸门;行情回补到 2005;手册补全流程
本轮会话的三项正确性改造(均为「不报错、只让结果静默错」的类型):

1) 修复 stock_daily 量价单位前后不一致
   - 现象:2015-2019 存 Tushare 原始单位(手/千元),2020 起存(股/元),2019 同日混合;
     而流动性阈值按「元」配置 → 早年门槛实际是「日均成交额 ≥ 200 亿元」,
     把 2015-2019 的股票池整体清空(实测 2016/2017/2018 各选出 0 只)。
   - 修复:写入端 sync/price.py 统一换算;读取端 units.normalize_ohlcv_units
     按行判定并幂等换算(price_history / avg_amount 都走它);
     审计新增 UNIT-OHLCV 防回归。
   - 效果:2016/2017/2018 的股票池变为 7/11/13 只。

2) 未来函数守卫(单次回测)
   - 股票池自带 asof:若晚于回测起点即**拒绝执行**(原先静默冻结套用),
     与 walk-forward 已有的拒绝理由一致;确需复现加 --allow-lookahead-universe,
     偏差写入 unimplemented_json。

3) 新增实时(PIT)个股画像闸门
   - profile/pit.py:每个决策日按当时可见数据重算过去 5 年画像,
     惰性(仅买入条件已触发的标的)、面板按 asof 缓存、
     规则不含财务指标时不查财报表;被剔除时产出 REJECT + 逐规则留痕。
   - 指标定义复用 ProfileBuilder._profile_one(与批量画像逐值等价的回归测试)。
   - profile/coverage.py:窗口覆盖率(按交易日历的真实开市天数),
     策略新增 entry.profile_gate.min_window_coverage(默认 0,不改变既有行为)。
   - core/metrics.py:闸门可用指标的唯一定义(配置期即校验,避免写错指标名静默失效)。

4) 行情回补到 2005(使 5/8/10 年窗口真正完整)
   - stock_daily / adjust_factor / daily_basic 补到 2005-01-04;
     hd_suspend / hd_limit 补到 2010-01-04。
   - 5 年窗口覆盖率:2018-05-18 由 67.0% → 99.1%,2016-12-30 由 39.8% → 99.0%;
     残差经逐日与 hd_suspend 交叉核实为真实停牌(16/16 命中)。
   - 审计 G2/G3 与断点续传原先用固定阈值(2000 / 1500 只),
     会把 2005-2009 的正常数据误判为异常 —— 改为按「当年应有上市股票数」成比例判定。
   - 节流修正:daily/adj_factor/daily_basic 限频 480 → 170(实测该 token 约 196/min 即被拒)。

5) 自我声明如实化
   - 原先「约束未生效」由「过滤后集合为空」判定,会把「这批股票恰好没停牌」
     误报成「hd_suspend 无数据」;改为按表级判定。
   - 补齐此前静默的「配置承诺但未实现」项:suspended_rule/limit_up_down_rule 的 defer、
     cash_mode=reinvest/reinvest_rule、handle_rights_issue、signal_to_execution、
     max_volume_pct、liquidity_limit_pct_adv —— 全部写入 unimplemented_json。

6) 手册:新增 §0「全流程操作(选股 → 画像 → 回测)」置于最前
   - 逐步说明「命令做了什么、数据从哪来、落了哪些库、有哪些坑」;
     含实时画像闸门 9 问 9 答、未来函数守卫表、成交与成本口径、验证 SQL。
   - 修正旧 §2.4 漏传 --universe-run(选了池子却没用于回测);
     修正两处声称「停牌顺延」「分红再投资」已实现的相反表述。

测试:403 项全部通过(含新增 test_units.py、test_profile_pit.py、
未实现声明诚实性测试、行序无关性回归测试)。

注意:本提交中 docs/*、README.md、src/hdiv/web/service.py 除本轮修改外,
也含此前遗留的未提交改动(无法按文件切分)。
2026-10-04 12:47:17 +08:00

940 lines
34 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 SyncConfig(StrictModel):
"""行情同步的「完整性」判定口径(决定断点续传会不会重拉整段历史)。
一个交易日被视为**已完整同步**,要求当日股票数 ≥
``max(min_symbols_floor, min_symbols_ratio × 当年应有上市股票数)``。
**为什么不能只用一个绝对阈值**:A 股 2005 年只有约 1,350 只股票,
2010 年约 1,700 只。若固定要求 1,500 只,2005-2009 的**每一个交易日**
都会被判成「未完成」,于是断点续传完全失效 —— 回补中断一次就要从
第一天重新拉,且每次重跑都会把整段早年历史再拉一遍。
"""
min_symbols_floor: int = 200
min_symbols_ratio: float = 0.6
@model_validator(mode="after")
def _check(self) -> SyncConfig:
if self.min_symbols_floor < 1:
raise SchemaValidationError("sync.min_symbols_floor 必须为正")
if not 0.0 < self.min_symbols_ratio <= 1.0:
raise SchemaValidationError("sync.min_symbols_ratio 必须落在 (0, 1]")
return self
class DataSourceConfig(StrictModel):
version: int = 1
database: DatabaseConfig
tushare: TushareConfig = Field(default_factory=TushareConfig)
paths: PathsConfig = Field(default_factory=PathsConfig)
sync: SyncConfig = Field(default_factory=SyncConfig)
# ---------------------------------------------------------------------------
# universe.yml
# ---------------------------------------------------------------------------
class EvaluationConfig(StrictModel):
mode: Literal["point_in_time", "latest"] = "point_in_time"
asof: date | Literal["latest"] = "latest"
class MarketFilterConfig(StrictModel):
exchanges: list[str] = Field(default_factory=lambda: ["SZSE", "SSE"])
markets: list[str] = Field(default_factory=lambda: ["主板", "创业板", "科创板"])
min_listing_years: int = 10
min_market_cap: float | None = 50_000_000_000.0
min_float_market_cap: float | None = None
min_avg_amount_20d: float | None = 20_000_000.0
require_trading_on_asof: bool = True
class RiskFilterConfig(StrictModel):
exclude_st: bool = True
st_lookback_from_history: bool = True
exclude_delisting: bool = True
exclude_suspended: bool = True
exclude_negative_equity: bool = True
max_debt_to_assets: float | None = 0.80
max_pledge_ratio: float | None = None
exclude_major_litigation: bool = False
exclude_goodwill_anomaly: bool = False
class DividendFilterConfig(StrictModel):
yield_source: Literal["computed", "dv_ttm"] = "computed"
min_dividend_yield: float | None = 0.03
min_continuous_years: int = 5
continuous_rule: Literal["cash_div_tax_gt_0"] = "cash_div_tax_gt_0"
window_years: int = 6
min_dividend_years_in_window: int = 5
max_payout_ratio: float | None = 1.00
min_payout_ratio: float | None = None
min_dps_cagr_5y: float | None = None
require_positive_fcf: bool = True
min_fcf_dividend_cover: float | None = 1.0
# 验证所需数据缺失时的处理:pass=放行(宽松)/ reject=淘汰(严格)
on_missing_data: Literal["pass", "reject"] = "pass"
@model_validator(mode="after")
def _check_years(self) -> DividendFilterConfig:
if self.min_continuous_years > self.window_years:
raise SchemaValidationError(
f"dividend.min_continuous_years({self.min_continuous_years}) "
f"不能大于 window_years({self.window_years})"
)
return self
class QualityFilterConfig(StrictModel):
min_roe_5y_avg: float | None = 0.08
min_roic_5y_avg: float | None = None
min_gross_margin: float | None = None
min_net_margin: float | None = None
min_ocf_to_profit: float | None = 0.60
max_debt_to_assets: float | None = None
min_interest_cover: float | None = None
exempt_industries: list[str] = Field(default_factory=list)
pit_rule: Literal["announce_date_le_asof"] = "announce_date_le_asof"
report_basis: Literal["latest_announced", "ttm"] = "latest_announced"
class UniverseOutputConfig(StrictModel):
min_members: int = 10
max_members: int | None = 200
warn_on_empty: bool = True
class IndustryExemptionsConfig(StrictModel):
"""按行业豁免特定阈值(金融股口径与实业不同)。"""
leverage: list[str] = Field(default_factory=lambda: ["银行", "保险", "证券", "信托"])
free_cash_flow: list[str] = Field(default_factory=lambda: ["银行", "保险", "证券", "信托"])
class UniverseConfig(StrictModel):
version: int = 1
name: str
evaluation: EvaluationConfig = Field(default_factory=EvaluationConfig)
industry_exemptions: IndustryExemptionsConfig = Field(
default_factory=IndustryExemptionsConfig
)
market: MarketFilterConfig = Field(default_factory=MarketFilterConfig)
risk: RiskFilterConfig = Field(default_factory=RiskFilterConfig)
dividend: DividendFilterConfig = Field(default_factory=DividendFilterConfig)
quality: QualityFilterConfig = Field(default_factory=QualityFilterConfig)
output: UniverseOutputConfig = Field(default_factory=UniverseOutputConfig)
# ---------------------------------------------------------------------------
# profile.yml
# ---------------------------------------------------------------------------
class DividendYieldVolatilityConfig(StrictModel):
daily_window: int = 250
monthly_window: int = 60
quarterly_window: int = 20
annual_window: int = 10
min_obs: int = 20
class ScoreAnchorPercentile(StrictModel):
best_percentile: float
worst_percentile: float
class ScoreAnchorAbs(StrictModel):
metric: str
best_abs: float
worst_abs: float
class ScoreAnchorQuality(StrictModel):
metric: str
best_abs: float
worst_abs: float
class ScoreAnchorsConfig(StrictModel):
dividend_yield: ScoreAnchorPercentile
valuation: ScoreAnchorPercentile
financial_quality: ScoreAnchorQuality
balance_sheet: ScoreAnchorQuality
dividend_quality: ScoreAnchorQuality
class SafetyMarginConfig(StrictModel):
mode: Literal["separate", "composite"] = "separate"
weights: dict[str, float]
score_bounds: tuple[float, float] = (0.0, 100.0)
score_anchors: ScoreAnchorsConfig | None = None
@model_validator(mode="after")
def _check_weights(self) -> SafetyMarginConfig:
if self.mode == "composite":
total = sum(self.weights.values())
if abs(total - 1.0) > 1e-6:
raise SchemaValidationError(
f"safety_margin.mode=composite 时 weights 之和须为 1.0,当前为 {total}"
)
if self.score_bounds[0] >= self.score_bounds[1]:
raise SchemaValidationError("score_bounds 必须满足 [min, max] 且 min < max")
return self
class MetricsConfig(StrictModel):
valuation: list[str] = Field(default_factory=list)
financial: list[str] = Field(default_factory=list)
dividend: list[str] = Field(default_factory=list)
risk: list[str] = Field(default_factory=list)
return_: list[str] = Field(default_factory=list, alias="return")
model_config = ConfigDict(extra="forbid", populate_by_name=True)
def all_codes(self) -> list[str]:
seen: list[str] = []
for group in (self.valuation, self.financial, self.dividend, self.risk, self.return_):
for code in group:
if code not in seen:
seen.append(code)
return seen
@model_validator(mode="after")
def _check_duplicates(self) -> MetricsConfig:
flat = self.valuation + self.financial + self.dividend + self.risk + self.return_
if len(flat) != len(set(flat)):
dupes = {c for c in flat if flat.count(c) > 1}
raise SchemaValidationError(f"profile.metrics 存在重复指标:{sorted(dupes)}")
return self
class SufficiencyConfig(StrictModel):
min_history_years_dividend: int = 5
min_history_years_price: int = 3
min_dividend_records: int = 5
class TtmDividendConfig(StrictModel):
window_days: int = 365
grace_days: int = 45
# 是否消除除权间隔不规整造成的毛刺(重叠虚高 / 断档虚低)
smooth_spikes: bool = True
@model_validator(mode="after")
def _check(self) -> TtmDividendConfig:
if self.window_days <= 0:
raise SchemaValidationError("ttm_dividend.window_days 必须为正")
if self.grace_days < 0:
raise SchemaValidationError("ttm_dividend.grace_days 不能为负")
return self
class ProfileConfig(StrictModel):
version: int = 1
name: str = "default_profile"
windows_years: list[int] = Field(default_factory=lambda: [5, 8, 10])
min_obs_days: int = 500
series_max_points: int = 1500
series_metrics: list[str] = Field(
default_factory=lambda: ["dv_yield", "pe_ttm", "pb", "close", "drawdown"]
)
ttm_dividend: TtmDividendConfig = Field(default_factory=TtmDividendConfig)
metrics: MetricsConfig = Field(default_factory=MetricsConfig)
percentiles: list[float] = Field(default_factory=lambda: [10, 25, 50, 75, 90])
dividend_yield_volatility: DividendYieldVolatilityConfig = Field(
default_factory=DividendYieldVolatilityConfig
)
safety_margin: SafetyMarginConfig
sufficiency: SufficiencyConfig = Field(default_factory=SufficiencyConfig)
@model_validator(mode="after")
def _check(self) -> ProfileConfig:
if not self.percentiles:
raise SchemaValidationError("profile.percentiles 不能为空")
if self.percentiles != sorted(self.percentiles):
raise SchemaValidationError("profile.percentiles 必须升序")
for p in self.percentiles:
if not 0 <= p <= 100:
raise SchemaValidationError(f"分位数越界:{p}")
if sorted(self.windows_years) != self.windows_years or not self.windows_years:
raise SchemaValidationError("profile.windows_years 必须为非空升序")
if self.min_obs_days <= 0:
raise SchemaValidationError("profile.min_obs_days 必须为正")
return self
# ---------------------------------------------------------------------------
# cost.yml
# ---------------------------------------------------------------------------
class CommissionConfig(StrictModel):
rate: float = 0.00025
min: float = 5.0
side: Literal["both", "buy", "sell"] = "both"
class SideRateConfig(StrictModel):
rate: float
side: Literal["both", "buy", "sell"] = "both"
class SlippageConfig(StrictModel):
mode: Literal["bps", "fixed", "tick"] = "bps"
value: float = 10.0
class DividendTaxConfig(StrictModel):
enabled: bool = True
holding_based: bool = True
rates: dict[str, float] = Field(default_factory=dict)
class CostConfig(StrictModel):
version: int = 1
model: str = "a_share_default"
commission: CommissionConfig = Field(default_factory=CommissionConfig)
stamp_duty: SideRateConfig = Field(
default_factory=lambda: SideRateConfig(rate=0.0005, side="sell")
)
transfer_fee: SideRateConfig = Field(
default_factory=lambda: SideRateConfig(rate=0.00001, side="both")
)
slippage: SlippageConfig = Field(default_factory=SlippageConfig)
dividend_tax: DividendTaxConfig = Field(default_factory=DividendTaxConfig)
# ---------------------------------------------------------------------------
# backtest.yml
# ---------------------------------------------------------------------------
class CapitalConfig(StrictModel):
initial: float = 1_000_000.0
currency: str = "CNY"
@model_validator(mode="after")
def _positive(self) -> CapitalConfig:
if self.initial <= 0:
raise SchemaValidationError("capital.initial 必须为正")
return self
class PeriodConfig(StrictModel):
start: date
end: DateOrLatest = "latest"
@model_validator(mode="after")
def _order(self) -> PeriodConfig:
if isinstance(self.end, date) and self.end <= self.start:
raise SchemaValidationError(
f"period.end({self.end}) 必须晚于 period.start({self.start})"
)
return self
class WalkForwardConfig(StrictModel):
enabled: bool = True
scheme: Literal["rolling", "expanding"] = "rolling"
train_years: int = 5
test_years: int = 1
step_months: int = 12
min_train_years: int = 5
freeze_params_in_test: bool = True
@model_validator(mode="after")
def _check(self) -> WalkForwardConfig:
if self.train_years < self.min_train_years:
raise SchemaValidationError(
f"walk_forward.train_years({self.train_years}) "
f"不能小于 min_train_years({self.min_train_years})"
)
if self.train_years <= 0 or self.test_years <= 0:
raise SchemaValidationError("walk_forward 的 train_years / test_years 必须为正")
if self.step_months <= 0:
raise SchemaValidationError("walk_forward.step_months 必须为正")
if not self.freeze_params_in_test:
raise SchemaValidationError(
"walk_forward.freeze_params_in_test 必须为 true"
"(plan.md §25:测试阶段禁止重新调整策略参数)"
)
return self
class DividendHandlingConfig(StrictModel):
cash_mode: Literal["reinvest", "hold", "cash_out"] = "reinvest"
reinvest_rule: Literal["same_stock_next_open", "portfolio_rebalance"] = (
"same_stock_next_open"
)
apply_dividend_tax: bool = True
handle_stock_dividend: bool = True
handle_rights_issue: bool = True
class BenchmarkConfig(StrictModel):
code: str
name: str
class FillConfig(StrictModel):
price: Literal["next_open", "next_close", "same_close"] = "next_open"
partial_fill: bool = False
max_volume_pct: float = 0.05
limit_up_down_rule: Literal["skip", "defer"] = "skip"
suspended_rule: Literal["skip", "defer"] = "defer"
class ReproConfig(StrictModel):
record_code_version: bool = True
record_data_version: bool = True
seed: int = 42
class ScheduleConfig(StrictModel):
signal_frequency_months: int = 1
universe_refresh_months: int = 12
@model_validator(mode="after")
def _check(self) -> ScheduleConfig:
for n, v in (("signal_frequency_months", self.signal_frequency_months),
("universe_refresh_months", self.universe_refresh_months)):
if v <= 0:
raise SchemaValidationError(f"schedule.{n} 必须为正")
return self
class PercentileReferenceConfig(StrictModel):
mode: Literal["rolling", "frozen"] = "rolling"
lookback_years: int = 5
min_observations: int = 250
@model_validator(mode="after")
def _check(self) -> PercentileReferenceConfig:
if self.lookback_years <= 0:
raise SchemaValidationError("percentile_reference.lookback_years 必须为正")
return self
class BacktestConfig(StrictModel):
version: int = 1
capital: CapitalConfig = Field(default_factory=CapitalConfig)
period: PeriodConfig
schedule: ScheduleConfig = Field(default_factory=ScheduleConfig)
percentile_reference: PercentileReferenceConfig = Field(
default_factory=PercentileReferenceConfig
)
walk_forward: WalkForwardConfig = Field(default_factory=WalkForwardConfig)
dividend: DividendHandlingConfig = Field(default_factory=DividendHandlingConfig)
benchmark: list[BenchmarkConfig] = Field(default_factory=list)
risk_free_rate: float = 0.02
fill: FillConfig = Field(default_factory=FillConfig)
reproducibility: ReproConfig = Field(default_factory=ReproConfig)
@model_validator(mode="after")
def _check_benchmarks(self) -> BacktestConfig:
codes = [b.code for b in self.benchmark]
if len(codes) != len(set(codes)):
raise SchemaValidationError(f"benchmark 存在重复代码:{codes}")
return self
# ---------------------------------------------------------------------------
# report.yml
# ---------------------------------------------------------------------------
class ChartsConfig(StrictModel):
kline_signals: bool = True
equity_curve: bool = True
drawdown: bool = True
yearly_returns: bool = True
monthly_heatmap: bool = True
rolling_metrics: bool = True
sector_exposure: bool = True
position_weights: bool = True
yield_percentile: bool = True
yield_histogram: bool = True
financial_trend: bool = True
sensitivity_heatmap: bool = True
walkforward_compare: bool = True
class NamingConfig(StrictModel):
index: str = "index.html"
audit: str = "data_audit_{date}.html"
universe: str = "universe_{asof}.html"
profile: str = "profile_{symbol}_{asof}.html"
backtest: str = "backtest_{run_id}.html"
walkforward: str = "walkforward_{wf_id}.html"
sensitivity: str = "sensitivity_{sens_id}.html"
class DecimalsConfig(StrictModel):
ratio: int = 4
money: int = 2
price: int = 2
class LayoutConfig(StrictModel):
max_width: int = 1440
table_page_size: int = 100
kline_default_days: int = 1500
normalize_nav: bool = True
decimals: DecimalsConfig = Field(default_factory=DecimalsConfig)
class ReportConfig(StrictModel):
version: int = 1
templates_dir: str = "templates"
output_dir: str = "output"
assets_dir: str = "assets"
theme: Literal["light", "dark", "auto"] = "light"
offline_assets: bool = True
asset_mode: Literal["shared", "inline"] = "shared"
include_sql_provenance: bool = True
charts: ChartsConfig = Field(default_factory=ChartsConfig)
naming: NamingConfig = Field(default_factory=NamingConfig)
layout: LayoutConfig = Field(default_factory=LayoutConfig)
# ---------------------------------------------------------------------------
# strategy/*.yml
# ---------------------------------------------------------------------------
class StrategyMeta(StrictModel):
id: str
name: str
version: str
status: Literal[
"DRAFT", "RESEARCH", "BACKTEST", "VALIDATED", "PAPER", "LIVE", "ARCHIVED"
] = "DRAFT"
description: str = ""
class UniverseRef(StrictModel):
ref: str = "config/universe.yml"
override: dict[str, Any] = Field(default_factory=dict)
class ScaleStep(StrictModel):
percentile: float
weight: float
@model_validator(mode="after")
def _check(self) -> ScaleStep:
if not 0 <= self.percentile <= 100:
raise SchemaValidationError(f"scale 步骤的分位数越界:{self.percentile}")
if not 0.0 <= self.weight <= 1.0:
raise SchemaValidationError(f"scale 步骤的权重越界:{self.weight}")
return self
class ProfileGateRule(StrictModel):
"""一条实时画像闸门规则。
语义:``<metric> 的 <stat> <op> <value>`` 必须成立,否则不买。
指标名与分位可用性在**配置期**校验 —— 写错一个指标名若拖到运行时,
只会得到「无法验证 → 保守不买」,表现为策略再也不交易,极难定位。
"""
metric: str
stat: Literal["current_value", "current_percentile"] = "current_value"
op: Literal[">=", "<=", ">", "<"] = ">="
value: float
@model_validator(mode="after")
def _check(self) -> ProfileGateRule:
from hdiv.core.metrics import GATE_METRICS, PERCENTILE_METRICS
if self.metric not in GATE_METRICS:
head = self.metric.split("_")[0]
near = sorted(m for m in GATE_METRICS if head and head in m)
raise SchemaValidationError(
f"profile_gate 规则引用了未知指标 {self.metric!r}。"
f"可选指标见 hdiv/core/metrics.py;相近的有 {near[:6]}"
)
if self.stat == "current_percentile" and self.metric not in PERCENTILE_METRICS:
raise SchemaValidationError(
f"{self.metric} 是标量指标,没有历史分位,不能用 "
f"stat=current_percentile。有分位的指标:{sorted(PERCENTILE_METRICS)}"
)
return self
class ProfileGateConfig(StrictModel):
"""实时(PIT)个股画像闸门。
在每个决策日、**买入条件已经触发之后**,用「当时可见的数据」重算画像,
不通过的票直接剔除。跨股票共享的面板按时点缓存,代价与「触发次数」成正比,
而不是与「回测区间 × 股票数」成正比。
"""
#: 是否启用。关闭时回测行为与启用前完全一致(可用于复现历史结果)
enabled: bool = False
#: 画像统计窗口(年)。0 = 全历史;其余必须是 config/profile.yml 的 windows_years 之一
window_years: int = 5
#: 数据缺失/样本不足(无法验证)时:reject = 保守不买,pass = 放行
on_unverifiable: Literal["reject", "pass"] = "reject"
#: 窗口**实际覆盖率**下限(1.0 = 名义 5 年就必须真有 5 年数据)。
#: 0 = 不因覆盖率淘汰(默认,保持改造前行为)。
#:
#: 为什么需要它:`window_slice(asof, 5)` 只是「把已有数据切成最近 5 年」,
#: 数据起点晚于窗口左端时窗口会被静默截短 —— 实测 2018-05-18 的「5 年」
#: 窗口只有 3.4 年(817/1215 个交易日,67%),而画像仍报 OK。
min_window_coverage: float = 0.0
rules: list[ProfileGateRule] = Field(default_factory=list)
@model_validator(mode="after")
def _check(self) -> ProfileGateConfig:
if self.window_years < 0:
raise SchemaValidationError("profile_gate.window_years 不能为负")
if not 0.0 <= self.min_window_coverage <= 1.0:
raise SchemaValidationError(
f"profile_gate.min_window_coverage 必须落在 [0, 1],"
f"当前 {self.min_window_coverage}"
)
if self.enabled and not self.rules:
raise SchemaValidationError(
"profile_gate.enabled=true 但 rules 为空 —— 空闸门等于每次都要"
"算一遍画像再无条件放行。请补齐规则,或把 enabled 设为 false。"
)
return self
class EntryConfig(StrictModel):
yield_percentile: float
require_universe_pass: bool = True
require_risk_pass: bool = True
scale_in: list[ScaleStep] = Field(default_factory=list)
profile_gate: ProfileGateConfig = Field(default_factory=ProfileGateConfig)
@model_validator(mode="after")
def _check(self) -> EntryConfig:
if not 0 <= self.yield_percentile <= 100:
raise SchemaValidationError(f"entry.yield_percentile 越界:{self.yield_percentile}")
if self.scale_in:
for a, b in zip(self.scale_in, self.scale_in[1:], strict=False):
if b.percentile <= a.percentile:
raise SchemaValidationError("entry.scale_in 必须按 percentile 严格升序")
if b.weight < a.weight:
raise SchemaValidationError("entry.scale_in 的 weight 必须随分位单调不减")
if self.scale_in[0].percentile != self.yield_percentile:
raise SchemaValidationError(
f"entry.scale_in 首个分位({self.scale_in[0].percentile}) "
f"必须等于 entry.yield_percentile({self.yield_percentile})"
)
return self
class ExitConfig(StrictModel):
yield_percentile: float
scale_out: list[ScaleStep] = Field(default_factory=list)
stop_loss_pct: float | None = None
max_holding_days: int | None = None
@model_validator(mode="after")
def _check(self) -> ExitConfig:
if not 0 <= self.yield_percentile <= 100:
raise SchemaValidationError(f"exit.yield_percentile 越界:{self.yield_percentile}")
if self.scale_out:
for a, b in zip(self.scale_out, self.scale_out[1:], strict=False):
if b.percentile >= a.percentile:
raise SchemaValidationError("exit.scale_out 必须按 percentile 严格降序")
if b.weight > a.weight:
raise SchemaValidationError("exit.scale_out 的 weight 必须随分位单调不减")
if self.scale_out[-1].percentile != self.yield_percentile:
raise SchemaValidationError(
f"exit.scale_out 末个分位({self.scale_out[-1].percentile}) "
f"必须等于 exit.yield_percentile({self.yield_percentile})"
)
if abs(self.scale_out[-1].weight) > 1e-9:
raise SchemaValidationError("exit.scale_out 末步 weight 必须为 0(清仓)")
return self
class PositionConfig(StrictModel):
max_position: float = 0.10
min_position: float = 0.01
sector_max_position: float = 0.25
max_holdings: int = 20
weight_scheme: Literal["equal", "score", "inverse_vol"] = "equal"
@model_validator(mode="after")
def _check(self) -> PositionConfig:
for name, v in (
("max_position", self.max_position),
("min_position", self.min_position),
("sector_max_position", self.sector_max_position),
):
if not 0.0 < v <= 1.0:
raise SchemaValidationError(f"position.{name} 必须落在 (0, 1],当前 {v}")
if self.min_position > self.max_position:
raise SchemaValidationError("position.min_position 不能大于 max_position")
if self.max_holdings <= 0:
raise SchemaValidationError("position.max_holdings 必须为正")
if self.max_position * self.max_holdings < 1.0 - 1e-9:
# 不阻断,但这是常见配置错误:单股上限 × 持仓数 < 1 会导致仓位无法打满
pass
return self
class RiskConfig(StrictModel):
max_portfolio_drawdown: float | None = 0.20
max_single_drawdown: float | None = 0.25
liquidity_limit_pct_adv: float | None = 0.05
class ExecutionConfig(StrictModel):
signal_to_execution: Literal["next_open", "next_close", "same_close"] = "next_open"
limit_up_down_rule: Literal["skip", "defer"] = "skip"
suspended_rule: Literal["skip", "defer"] = "defer"
class CostRef(StrictModel):
ref: str = "config/cost.yml"
class StrategyConfig(StrictModel):
strategy: StrategyMeta
universe: UniverseRef = Field(default_factory=UniverseRef)
entry: EntryConfig
exit: ExitConfig
position: PositionConfig = Field(default_factory=PositionConfig)
risk: RiskConfig = Field(default_factory=RiskConfig)
execution: ExecutionConfig = Field(default_factory=ExecutionConfig)
cost: CostRef = Field(default_factory=CostRef)
@model_validator(mode="after")
def _check_entry_exit(self) -> StrategyConfig:
if self.entry.yield_percentile <= self.exit.yield_percentile:
raise SchemaValidationError(
f"entry.yield_percentile({self.entry.yield_percentile}) 必须大于 "
f"exit.yield_percentile({self.exit.yield_percentile})"
)
return self
def param_map(self) -> dict[str, Any]:
"""扁平化参数路径 → 值,写入 hd_strategy_param,便于跨版本 diff 与敏感性扫描。"""
out: dict[str, Any] = {}
def walk(prefix: str, obj: Any) -> None:
if isinstance(obj, BaseModel):
for name in type(obj).model_fields:
walk(f"{prefix}.{name}" if prefix else name, getattr(obj, name))
elif isinstance(obj, list):
for i, item in enumerate(obj):
walk(f"{prefix}[{i}]", item)
elif isinstance(obj, dict):
for k, v in obj.items():
walk(f"{prefix}.{k}" if prefix else str(k), v)
else:
out[prefix] = obj
walk("", self)
return out
# ---------------------------------------------------------------------------
# 加载与哈希
# ---------------------------------------------------------------------------
_MODEL_BY_NAME: dict[str, type[StrictModel]] = {
"datasource": DataSourceConfig,
"universe": UniverseConfig,
"profile": ProfileConfig,
"cost": CostConfig,
"backtest": BacktestConfig,
"report": ReportConfig,
}
def resolve_strategy_path(rel: str | Path) -> Path:
"""策略路径解析,按以下顺序尝试:
1. 绝对路径 → 原样
2. 项目根相对(如 ``config/strategy/high_dividend_v1.yml``)
3. ``config/`` 相对(如 ``strategy/high_dividend_v1.yml``)
4. ``config/strategy/`` 下的裸文件名(如 ``high_dividend_v1.yml``)
"""
p = Path(rel)
if p.is_absolute():
return p
candidates = [
project_root() / p,
config_dir() / p,
config_dir() / "strategy" / p,
]
for cand in candidates:
if cand.is_file():
return cand
raise ConfigNotFound(
"策略配置不存在,已尝试:\n " + "\n ".join(str(c) for c in candidates)
)
def _read_yaml(path: Path) -> dict[str, Any]:
if not path.is_file():
raise ConfigNotFound(f"配置文件不存在:{path}")
try:
raw = yaml.safe_load(path.read_text(encoding="utf-8"))
except yaml.YAMLError as exc: # pragma: no cover - 由 YAML 库决定
raise ConfigError(f"YAML 解析失败:{path}\n{exc}") from exc
if raw is None:
raise ConfigError(f"配置文件为空:{path}")
if not isinstance(raw, dict):
raise ConfigError(f"配置文件顶层必须是映射(dict):{path}")
return raw
def load_config(name: str, path: str | Path | None = None) -> Any:
"""按配置名加载并校验。
name ∈ {datasource, universe, profile, cost, backtest, report}
或 'strategy:<相对路径>' 形式加载策略文件。
"""
if name.startswith("strategy:"):
rel = name.split(":", 1)[1]
p = resolve_strategy_path(rel)
raw = _read_yaml(p)
try:
return StrategyConfig.model_validate(raw)
except Exception as exc:
raise ConfigError(f"策略配置校验失败:{p}\n{exc}") from exc
if name not in _MODEL_BY_NAME:
raise ConfigError(f"未知配置名:{name}(可用:{sorted(_MODEL_BY_NAME)})")
p = Path(path) if path else config_dir() / f"{name}.yml"
raw = _read_yaml(p)
try:
return _MODEL_BY_NAME[name].model_validate(raw)
except Exception as exc:
raise ConfigError(f"配置校验失败:{p}\n{exc}") from exc
@lru_cache(maxsize=32)
def _cached_config(name: str) -> Any:
return load_config(name)
def get_config(name: str) -> Any:
"""带缓存的配置读取(同进程内只解析一次)。"""
return _cached_config(name)
def config_hash(cfg: BaseModel | Any) -> str:
"""配置指纹:用于 run 记录与可复现性校验。"""
if isinstance(cfg, BaseModel):
payload = cfg.model_dump(mode="json", exclude_none=False)
elif isinstance(cfg, dict):
payload = cfg
else:
payload = {"value": str(cfg)}
blob = json.dumps(payload, sort_keys=True, ensure_ascii=False, default=str)
return hashlib.sha256(blob.encode("utf-8")).hexdigest()[:32]
def load_all(names: tuple[str, ...] = ("datasource", "universe", "profile", "cost", "backtest", "report")) -> dict[str, Any]:
"""一次性加载全部主配置。"""
return {n: load_config(n) for n in names}