rebuild(stage-9): deliver deterministic intelligent screening

This commit is contained in:
leefer
2026-07-30 05:15:17 +08:00
parent 6cb52e864a
commit 158257ebb8
46 changed files with 7322 additions and 34 deletions
+134
View File
@@ -0,0 +1,134 @@
from __future__ import annotations
import json
import math
from functools import lru_cache
from pathlib import Path
from typing import Any
CONFIG_ROOT = Path(__file__).resolve().parents[3] / "config"
ALLOWED_OPERATORS = frozenset({">", ">=", "<", "<=", "==", "!=", "between", "in"})
class CatalogError(RuntimeError):
pass
@lru_cache(maxsize=1)
def factor_catalog() -> dict[str, Any]:
payload = _read("screener-factors.json")
factors = payload.get("factors")
groups = payload.get("groups")
if payload.get("schema_version") != 1 or not isinstance(factors, dict):
raise CatalogError("选股因子目录无效")
grouped = [field for values in (groups or {}).values() for field in values]
if len(grouped) != len(set(grouped)) or set(grouped) != set(factors):
raise CatalogError("选股因子分组与目录不一致")
return payload
@lru_cache(maxsize=1)
def strategy_catalog() -> tuple[dict[str, Any], ...]:
payload = _read("screener-strategies.json")
items = payload.get("strategies")
if payload.get("schema_version") != 1 or not isinstance(items, list):
raise CatalogError("选股策略目录无效")
identifiers: set[str] = set()
names: set[str] = set()
factors = set(factor_catalog()["factors"])
for item in items:
identifier = str(item.get("id") or "")
name = str(item.get("name") or "")
if not identifier or identifier in identifiers or not name or name in names:
raise CatalogError("选股策略标识或名称重复")
identifiers.add(identifier)
names.add(name)
validate_formula(item.get("formula"), factors)
if sum(item.get("kind") == "stage" for item in items) != 7:
raise CatalogError("阶段策略必须为7套")
if sum(item.get("kind") == "curated" for item in items) != 29:
raise CatalogError("精选策略必须为29套")
return tuple(items)
def strategy_by_id(identifier: str) -> dict[str, Any] | None:
return next((item for item in strategy_catalog() if item["id"] == identifier), None)
def validate_formula(formula: Any, factors: set[str] | None = None) -> dict[str, Any]:
if not isinstance(formula, dict):
raise CatalogError("选股公式必须是对象")
known = factors or set(factor_catalog()["factors"])
universe = formula.get("universe") or {}
listed_days = universe.get("listed_days_min", 120)
if not isinstance(listed_days, int) or not 0 <= listed_days <= 5000:
raise CatalogError("上市天数范围无效")
filters = formula.get("filters")
scores = formula.get("score")
if not isinstance(filters, list) or len(filters) > 20:
raise CatalogError("筛选条件必须为不超过20项的列表")
if not isinstance(scores, list) or not 1 <= len(scores) <= 12:
raise CatalogError("评分因子必须为1至12项")
for condition in filters:
if condition.get("field") not in known:
raise CatalogError(f"未知筛选因子:{condition.get('field')}")
if condition.get("op") not in ALLOWED_OPERATORS or "value" not in condition:
raise CatalogError("筛选运算符或比较值无效")
_validate_comparison(condition["op"], condition["value"])
total = 0.0
score_fields: set[str] = set()
for score in scores:
if score.get("field") not in known:
raise CatalogError(f"未知评分因子:{score.get('field')}")
field = str(score["field"])
if field in score_fields:
raise CatalogError("评分因子不能重复")
score_fields.add(field)
raw_weight = score.get("weight")
if not isinstance(raw_weight, (int, float)) or isinstance(raw_weight, bool):
raise CatalogError("评分权重必须是数值")
weight = float(raw_weight)
if weight <= 0 or score.get("direction", "desc") not in {"asc", "desc"}:
raise CatalogError("评分权重或方向无效")
total += weight
if abs(total - 1) > 0.000001:
raise CatalogError("评分权重总和必须为100%")
limit = formula.get("limit")
minimum = formula.get("min_score")
if not isinstance(limit, int) or not 1 <= limit <= 50:
raise CatalogError("输出数量必须为1至50")
if not isinstance(minimum, (int, float)) or not 0 <= minimum <= 1:
raise CatalogError("最低综合分必须在0至1之间")
return formula
def _validate_comparison(operator: str, value: Any) -> None:
if operator == "between":
if not isinstance(value, list) or len(value) != 2:
raise CatalogError("区间条件必须包含两个边界")
if not all(_comparable(item) for item in value) or value[0] > value[1]:
raise CatalogError("区间条件边界无效")
return
if operator == "in":
if not isinstance(value, list) or not 1 <= len(value) <= 20:
raise CatalogError("集合条件必须包含1至20个值")
if not all(_comparable(item) for item in value):
raise CatalogError("集合条件包含无效值")
return
if not _comparable(value):
raise CatalogError("比较值必须是有限数值或布尔值")
def _comparable(value: Any) -> bool:
return isinstance(value, bool) or (
isinstance(value, (int, float))
and not isinstance(value, bool)
and math.isfinite(float(value))
)
def _read(filename: str) -> dict[str, Any]:
try:
return json.loads((CONFIG_ROOT / filename).read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise CatalogError(f"无法读取{filename}") from exc