708 lines
32 KiB
Python
708 lines
32 KiB
Python
from __future__ import annotations
|
||
|
||
import copy
|
||
import json
|
||
import math
|
||
import statistics
|
||
from collections import defaultdict
|
||
from datetime import datetime, timedelta
|
||
from typing import Any
|
||
|
||
from database import ReviewDatabase
|
||
from sentiment_engine import build_sentiment_history, latest_contiguous_history
|
||
from tushare_client import TushareClient, TushareError
|
||
|
||
|
||
REGIMES = {
|
||
"ice": "冰点",
|
||
"repair": "修复",
|
||
"fermentation": "发酵",
|
||
"climax": "高潮",
|
||
"divergence": "分化",
|
||
"retreat": "退潮",
|
||
}
|
||
|
||
FACTOR_FIELDS = {
|
||
"pct_chg": "当日涨幅",
|
||
"return_5d": "5日涨幅",
|
||
"return_10d": "10日涨幅",
|
||
"above_ma20": "站上20日线",
|
||
"volume_ratio_5d": "5日量比",
|
||
"volatility_10d": "10日波动率",
|
||
"amount_billion": "成交额",
|
||
"turnover_rate": "换手率",
|
||
"circ_mv_billion": "流通市值",
|
||
"net_flow_million": "主力净流入",
|
||
"large_flow_million": "大单净流入",
|
||
"sector_strength": "板块强度",
|
||
"sector_limit_count": "板块涨停数",
|
||
"sector_up_count": "板块强势股数",
|
||
"relative_strength": "相对强度",
|
||
"limit_streak": "连板高度",
|
||
}
|
||
|
||
ALLOWED_OPERATORS = {">", ">=", "<", "<=", "==", "!=", "between", "in"}
|
||
|
||
|
||
BUILTIN_STRATEGIES = [
|
||
{
|
||
"name": "冰点抗跌先手",
|
||
"description": "寻找冰点中保持相对强度、低波动且有板块承接的个股,允许无结果。",
|
||
"regimes": ["ice"],
|
||
"formula": {
|
||
"universe": {"exclude_st": True, "listed_days_min": 120},
|
||
"filters": [
|
||
{"field": "pct_chg", "op": "between", "value": [-3, 7]},
|
||
{"field": "return_5d", "op": ">=", "value": -5},
|
||
{"field": "amount_billion", "op": ">=", "value": 1},
|
||
{"field": "volatility_10d", "op": "<=", "value": 7},
|
||
],
|
||
"score": [
|
||
{"field": "relative_strength", "weight": 0.30, "direction": "desc"},
|
||
{"field": "sector_strength", "weight": 0.25, "direction": "desc"},
|
||
{"field": "volume_ratio_5d", "weight": 0.20, "direction": "desc"},
|
||
{"field": "volatility_10d", "weight": 0.15, "direction": "asc"},
|
||
{"field": "amount_billion", "weight": 0.10, "direction": "desc"},
|
||
],
|
||
"limit": 12,
|
||
"min_score": 0.58,
|
||
},
|
||
},
|
||
{
|
||
"name": "修复先锋",
|
||
"description": "筛选率先站回趋势、温和放量并获得板块共振的修复前排。",
|
||
"regimes": ["repair"],
|
||
"formula": {
|
||
"universe": {"exclude_st": True, "listed_days_min": 120},
|
||
"filters": [
|
||
{"field": "pct_chg", "op": "between", "value": [1, 9.7]},
|
||
{"field": "return_5d", "op": ">", "value": 0},
|
||
{"field": "above_ma20", "op": "==", "value": 1},
|
||
{"field": "volume_ratio_5d", "op": ">=", "value": 1.05},
|
||
],
|
||
"score": [
|
||
{"field": "sector_strength", "weight": 0.28, "direction": "desc"},
|
||
{"field": "relative_strength", "weight": 0.24, "direction": "desc"},
|
||
{"field": "volume_ratio_5d", "weight": 0.18, "direction": "desc"},
|
||
{"field": "net_flow_million", "weight": 0.16, "direction": "desc"},
|
||
{"field": "amount_billion", "weight": 0.14, "direction": "desc"},
|
||
],
|
||
"limit": 15,
|
||
"min_score": 0.54,
|
||
},
|
||
},
|
||
{
|
||
"name": "主线发酵跟随",
|
||
"description": "在主线扩散期寻找趋势、成交承载和板块涨停梯队共同增强的个股。",
|
||
"regimes": ["fermentation"],
|
||
"formula": {
|
||
"universe": {"exclude_st": True, "listed_days_min": 120},
|
||
"filters": [
|
||
{"field": "pct_chg", "op": "between", "value": [0, 9.8]},
|
||
{"field": "return_5d", "op": ">=", "value": 3},
|
||
{"field": "above_ma20", "op": "==", "value": 1},
|
||
{"field": "amount_billion", "op": ">=", "value": 2},
|
||
],
|
||
"score": [
|
||
{"field": "sector_limit_count", "weight": 0.25, "direction": "desc"},
|
||
{"field": "sector_strength", "weight": 0.24, "direction": "desc"},
|
||
{"field": "return_10d", "weight": 0.20, "direction": "desc"},
|
||
{"field": "amount_billion", "weight": 0.16, "direction": "desc"},
|
||
{"field": "large_flow_million", "weight": 0.15, "direction": "desc"},
|
||
],
|
||
"limit": 15,
|
||
"min_score": 0.55,
|
||
},
|
||
},
|
||
{
|
||
"name": "高潮核心去后排",
|
||
"description": "高潮阶段只保留容量、趋势和辨识度较高的核心,降低后排跟风权重。",
|
||
"regimes": ["climax"],
|
||
"formula": {
|
||
"universe": {"exclude_st": True, "listed_days_min": 120},
|
||
"filters": [
|
||
{"field": "pct_chg", "op": "between", "value": [-2, 7]},
|
||
{"field": "return_10d", "op": ">=", "value": 5},
|
||
{"field": "above_ma20", "op": "==", "value": 1},
|
||
{"field": "amount_billion", "op": ">=", "value": 5},
|
||
],
|
||
"score": [
|
||
{"field": "amount_billion", "weight": 0.28, "direction": "desc"},
|
||
{"field": "sector_strength", "weight": 0.22, "direction": "desc"},
|
||
{"field": "relative_strength", "weight": 0.20, "direction": "desc"},
|
||
{"field": "volatility_10d", "weight": 0.15, "direction": "asc"},
|
||
{"field": "limit_streak", "weight": 0.15, "direction": "desc"},
|
||
],
|
||
"limit": 10,
|
||
"min_score": 0.62,
|
||
},
|
||
},
|
||
{
|
||
"name": "分化承接回流",
|
||
"description": "寻找分化中仍有趋势承接、板块强度和资金回流的核心候选。",
|
||
"regimes": ["divergence"],
|
||
"formula": {
|
||
"universe": {"exclude_st": True, "listed_days_min": 120},
|
||
"filters": [
|
||
{"field": "pct_chg", "op": "between", "value": [-3, 7]},
|
||
{"field": "return_5d", "op": ">", "value": 0},
|
||
{"field": "above_ma20", "op": "==", "value": 1},
|
||
{"field": "volume_ratio_5d", "op": "between", "value": [0.7, 3.5]},
|
||
],
|
||
"score": [
|
||
{"field": "relative_strength", "weight": 0.28, "direction": "desc"},
|
||
{"field": "sector_strength", "weight": 0.24, "direction": "desc"},
|
||
{"field": "net_flow_million", "weight": 0.20, "direction": "desc"},
|
||
{"field": "volatility_10d", "weight": 0.16, "direction": "asc"},
|
||
{"field": "amount_billion", "weight": 0.12, "direction": "desc"},
|
||
],
|
||
"limit": 12,
|
||
"min_score": 0.57,
|
||
},
|
||
},
|
||
{
|
||
"name": "退潮防守观察",
|
||
"description": "退潮期采用高门槛防守筛选,结果为空代表当前不宜主动出击。",
|
||
"regimes": ["retreat"],
|
||
"formula": {
|
||
"universe": {"exclude_st": True, "listed_days_min": 180},
|
||
"filters": [
|
||
{"field": "pct_chg", "op": "between", "value": [-2, 4]},
|
||
{"field": "return_5d", "op": ">=", "value": -2},
|
||
{"field": "above_ma20", "op": "==", "value": 1},
|
||
{"field": "volatility_10d", "op": "<=", "value": 4.5},
|
||
{"field": "amount_billion", "op": ">=", "value": 2},
|
||
],
|
||
"score": [
|
||
{"field": "volatility_10d", "weight": 0.30, "direction": "asc"},
|
||
{"field": "relative_strength", "weight": 0.25, "direction": "desc"},
|
||
{"field": "amount_billion", "weight": 0.20, "direction": "desc"},
|
||
{"field": "sector_strength", "weight": 0.15, "direction": "desc"},
|
||
{"field": "net_flow_million", "weight": 0.10, "direction": "desc"},
|
||
],
|
||
"limit": 8,
|
||
"min_score": 0.68,
|
||
},
|
||
},
|
||
]
|
||
|
||
|
||
class FactorDataService:
|
||
def __init__(self, database: ReviewDatabase, client: TushareClient) -> None:
|
||
self.database = database
|
||
self.client = client
|
||
|
||
def sync(self, requested_date: str, lookback: int = 45) -> dict[str, Any]:
|
||
trade_date, _ = self.client.resolve_trade_context(requested_date)
|
||
end = datetime.strptime(trade_date, "%Y%m%d")
|
||
start = (end - timedelta(days=max(100, lookback * 2 + 20))).strftime("%Y%m%d")
|
||
calendar = self.client.query(
|
||
"trade_cal",
|
||
{"exchange": "SSE", "start_date": start, "end_date": trade_date, "is_open": 1},
|
||
"cal_date,is_open",
|
||
)
|
||
dates = sorted(row["cal_date"] for row in calendar if row.get("is_open") == 1)[-lookback:]
|
||
existing = set(self.database.factor_dates(trade_date, lookback + 10))
|
||
dates_to_fetch = [value for value in dates if value not in existing or value == trade_date]
|
||
|
||
master = self.client.query(
|
||
"stock_basic",
|
||
{"list_status": "L"},
|
||
"ts_code,name,industry,market,list_date",
|
||
)
|
||
master_count = self.database.upsert_stock_master(master)
|
||
bar_count = 0
|
||
for current_date in dates_to_fetch:
|
||
rows = self.client.query(
|
||
"daily",
|
||
{"trade_date": current_date},
|
||
"ts_code,trade_date,open,high,low,close,pct_chg,vol,amount",
|
||
)
|
||
bar_count += self.database.upsert_daily_bars(rows)
|
||
|
||
indicators = self.client.query(
|
||
"daily_basic",
|
||
{"trade_date": trade_date},
|
||
"ts_code,trade_date,turnover_rate,volume_ratio,total_mv,circ_mv",
|
||
)
|
||
indicator_count = self.database.upsert_daily_indicators(indicators)
|
||
notices = []
|
||
try:
|
||
moneyflow = self.client.query(
|
||
"moneyflow",
|
||
{"trade_date": trade_date},
|
||
"ts_code,trade_date,buy_sm_amount,sell_sm_amount,buy_md_amount,sell_md_amount,"
|
||
"buy_lg_amount,sell_lg_amount,buy_elg_amount,sell_elg_amount,net_mf_amount",
|
||
)
|
||
moneyflow_count = self.database.upsert_moneyflow(moneyflow)
|
||
except TushareError as exc:
|
||
moneyflow_count = 0
|
||
notices.append(f"资金流接口不可用:{exc}")
|
||
|
||
return {
|
||
"trade_date": trade_date,
|
||
"calendar_dates": len(dates),
|
||
"fetched_dates": len(dates_to_fetch),
|
||
"stocks": master_count,
|
||
"bars": bar_count,
|
||
"indicators": indicator_count,
|
||
"moneyflow": moneyflow_count,
|
||
"notice": ";".join(notices),
|
||
}
|
||
|
||
|
||
class ScreenerEngine:
|
||
def __init__(self, database: ReviewDatabase) -> None:
|
||
self.database = database
|
||
|
||
def ensure_builtin_strategies(self) -> None:
|
||
existing = {item["name"] for item in self.database.list_screener_strategies() if item["builtin"]}
|
||
for strategy in BUILTIN_STRATEGIES:
|
||
if strategy["name"] not in existing:
|
||
self.database.save_screener_strategy(None, **strategy, builtin=True)
|
||
|
||
def detect_regime(self, trade_date: str) -> dict[str, Any]:
|
||
series = latest_contiguous_history(
|
||
build_sentiment_history(self.database.list_snapshot_payloads(trade_date, 240))
|
||
)
|
||
if not series:
|
||
return {
|
||
"id": "repair", "label": REGIMES["repair"], "confidence": 25,
|
||
"reason": "复盘快照不足,暂按中性修复处理。", "evidence": [], "history": [],
|
||
}
|
||
current = series[-1]
|
||
previous = series[-2] if len(series) > 1 else current
|
||
score = _number(current.get("score"))
|
||
previous_score = _number(previous.get("score"))
|
||
delta = score - previous_score
|
||
seal_rate = _number(current.get("seal_rate"))
|
||
limit_up = _number(current.get("limit_up_count"))
|
||
broken = _number(current.get("broken_count"))
|
||
regime = next(
|
||
(key for key, label in REGIMES.items() if label == current.get("phase")),
|
||
"divergence",
|
||
)
|
||
confidence = min(92, 45 + len(series[-8:]) * 5 + min(abs(delta), 12))
|
||
evidence = [
|
||
f"情绪温度 {score:.0f},较前一交易日 {delta:+.0f},{current.get('direction') or '持平'}",
|
||
f"封板率 {seal_rate:.1f}%",
|
||
f"涨停 {limit_up:.0f} 家,炸板 {broken:.0f} 家",
|
||
]
|
||
return {
|
||
"id": regime,
|
||
"label": REGIMES[regime],
|
||
"confidence": round(confidence),
|
||
"reason": _regime_reason(regime),
|
||
"evidence": evidence,
|
||
"history": [
|
||
{"trade_date": item["trade_date"], "score": _number(item.get("score"))}
|
||
for item in series[-8:]
|
||
],
|
||
}
|
||
|
||
def validate_formula(self, formula: dict[str, Any]) -> dict[str, Any]:
|
||
if not isinstance(formula, dict):
|
||
raise ValueError("选股公式必须是 JSON 对象。")
|
||
result = copy.deepcopy(formula)
|
||
universe = result.setdefault("universe", {})
|
||
universe["exclude_st"] = bool(universe.get("exclude_st", True))
|
||
universe["listed_days_min"] = max(0, min(5000, int(universe.get("listed_days_min", 120))))
|
||
filters = result.setdefault("filters", [])
|
||
if not isinstance(filters, list) or len(filters) > 20:
|
||
raise ValueError("筛选条件必须是列表,且不能超过 20 条。")
|
||
for condition in filters:
|
||
field = condition.get("field")
|
||
operator = condition.get("op")
|
||
if field not in FACTOR_FIELDS:
|
||
raise ValueError(f"不支持的选股因子:{field}")
|
||
if operator not in ALLOWED_OPERATORS:
|
||
raise ValueError(f"不支持的运算符:{operator}")
|
||
if "value" not in condition:
|
||
raise ValueError(f"因子 {field} 缺少比较值。")
|
||
scores = result.setdefault("score", [])
|
||
if not isinstance(scores, list) or not scores or len(scores) > 12:
|
||
raise ValueError("评分因子应为 1 至 12 条。")
|
||
for item in scores:
|
||
if item.get("field") not in FACTOR_FIELDS:
|
||
raise ValueError(f"不支持的评分因子:{item.get('field')}")
|
||
item["weight"] = float(item.get("weight", 0))
|
||
if item["weight"] <= 0 or item["weight"] > 1:
|
||
raise ValueError("评分权重必须大于 0 且不超过 1。")
|
||
if item.get("direction", "desc") not in {"asc", "desc"}:
|
||
raise ValueError("评分方向只能是 asc 或 desc。")
|
||
item["direction"] = item.get("direction", "desc")
|
||
result["limit"] = max(1, min(50, int(result.get("limit", 15))))
|
||
result["min_score"] = max(0, min(1, float(result.get("min_score", 0))))
|
||
return result
|
||
|
||
def screen(
|
||
self, user_id: int, trade_date: str, formula: dict[str, Any], regime: str,
|
||
strategy_name: str, run_backtest: bool = True,
|
||
realtime_snapshot: dict[str, Any] | None = None,
|
||
) -> dict[str, Any]:
|
||
formula = self.validate_formula(formula)
|
||
factors, actual_date = self.build_factors(trade_date, realtime_snapshot)
|
||
candidates = self.apply_formula(factors, formula, regime)
|
||
backtest = self.backtest(actual_date, formula) if run_backtest else None
|
||
if backtest and backtest["samples"] >= 20:
|
||
for candidate in candidates:
|
||
estimate = backtest["win_rate"] * 0.65 + candidate["score"] * 100 * 0.35
|
||
candidate["historical_probability"] = round(min(95, max(5, estimate)), 1)
|
||
candidate["probability_samples"] = backtest["samples"]
|
||
else:
|
||
for candidate in candidates:
|
||
candidate["historical_probability"] = None
|
||
candidate["probability_samples"] = backtest["samples"] if backtest else 0
|
||
result = {
|
||
"meta": {
|
||
"trade_date": _display_date(actual_date),
|
||
"regime": regime,
|
||
"regime_label": REGIMES.get(regime, regime),
|
||
"strategy_name": strategy_name,
|
||
"universe_count": len(factors),
|
||
"candidate_count": len(candidates),
|
||
"updated_at": datetime.now().astimezone().isoformat(timespec="seconds"),
|
||
"selection_source": (
|
||
"tushare_rt_k+history" if realtime_snapshot else "historical_eod"
|
||
),
|
||
"realtime": bool(realtime_snapshot),
|
||
"history_cutoff": (
|
||
str(realtime_snapshot.get("previous_trade_date") or "")
|
||
if realtime_snapshot else actual_date
|
||
),
|
||
"factor_freshness": {
|
||
"realtime": [
|
||
"价格", "涨跌幅", "成交量", "成交额", "换手率",
|
||
"均线位置", "5/10日动量", "板块强度",
|
||
] if realtime_snapshot else [],
|
||
"historical": ["历史波动率", "流通市值", "资金流", "回测"],
|
||
},
|
||
},
|
||
"formula": formula,
|
||
"candidates": candidates,
|
||
"backtest": backtest,
|
||
"disclaimer": "概率为历史条件估计,不代表未来收益;退潮或样本不足时允许无候选。",
|
||
}
|
||
run_id = self.database.save_screener_run(
|
||
user_id, actual_date, regime, strategy_name, formula, result
|
||
)
|
||
result["meta"]["run_id"] = run_id
|
||
return result
|
||
|
||
def build_factors(
|
||
self,
|
||
trade_date: str,
|
||
realtime_snapshot: dict[str, Any] | None = None,
|
||
) -> tuple[list[dict[str, Any]], str]:
|
||
data = self.database.load_factor_data(trade_date, 80)
|
||
dates = [value for value in data["dates"] if value <= trade_date]
|
||
if len(dates) < 21:
|
||
raise ValueError("历史行情不足 21 个交易日,请先同步因子数据。")
|
||
history_date = dates[-1]
|
||
realtime_map = {
|
||
str(row.get("ts_code") or ""): row
|
||
for row in (realtime_snapshot or {}).get("rows") or []
|
||
}
|
||
realtime_date = str((realtime_snapshot or {}).get("trade_date") or "")
|
||
use_realtime = bool(realtime_map and realtime_date == trade_date and history_date < trade_date)
|
||
actual_date = trade_date if use_realtime else history_date
|
||
master = {row["ts_code"]: row for row in data["master"]}
|
||
indicators = {row["ts_code"]: row for row in data["indicators"]}
|
||
moneyflow = {row["ts_code"]: row for row in data["moneyflow"]}
|
||
grouped: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||
for row in data["bars"]:
|
||
if row["trade_date"] <= history_date:
|
||
grouped[row["ts_code"]].append(row)
|
||
|
||
snapshot = self.database.get_snapshot(actual_date) or {}
|
||
limit_map: dict[str, tuple[str, int]] = {}
|
||
for key, status in (("limits", "涨停"), ("broken", "炸板"), ("down_limits", "跌停")):
|
||
for row in snapshot.get(key) or []:
|
||
limit_map[str(row.get("code"))] = (status, int(row.get("streak") or 0))
|
||
|
||
factors = []
|
||
current_day = datetime.strptime(actual_date, "%Y%m%d")
|
||
for ts_code, bars in grouped.items():
|
||
bars.sort(key=lambda item: item["trade_date"])
|
||
if len(bars) < 21 or bars[-1]["trade_date"] != history_date:
|
||
continue
|
||
info = master.get(ts_code)
|
||
if not info:
|
||
continue
|
||
historical_closes = [_number(item["close"]) for item in bars]
|
||
historical_volumes = [_number(item["vol"]) for item in bars]
|
||
realtime = realtime_map.get(ts_code) if use_realtime else None
|
||
current = realtime or bars[-1]
|
||
closes = historical_closes + ([_number(realtime["close"])] if realtime else [])
|
||
volumes = historical_volumes + ([_number(realtime["vol"])] if realtime else [])
|
||
if closes[-1] <= 0:
|
||
continue
|
||
returns_10 = [_number(item["pct_chg"]) for item in bars[-10:]]
|
||
if realtime:
|
||
returns_10 = returns_10[-9:] + [_number(realtime.get("pct_chg"))]
|
||
previous_volume = statistics.fmean(volumes[-6:-1]) if any(volumes[-6:-1]) else 0
|
||
indicator = indicators.get(ts_code, {})
|
||
flow = moneyflow.get(ts_code, {})
|
||
list_date = str(info.get("list_date") or "")
|
||
try:
|
||
listed_days = (current_day - datetime.strptime(list_date, "%Y%m%d")).days
|
||
except ValueError:
|
||
listed_days = 9999
|
||
code = str(info.get("code") or ts_code.split(".")[0])
|
||
status, streak = limit_map.get(code, ("", 0))
|
||
factors.append(
|
||
{
|
||
"code": code,
|
||
"ts_code": ts_code,
|
||
"name": info.get("name") or "--",
|
||
"sector": info.get("industry") or "其他",
|
||
"market": info.get("market") or "--",
|
||
"listed_days": listed_days,
|
||
"price": round(closes[-1], 2),
|
||
"pct_chg": round(_number(current["pct_chg"]), 2),
|
||
"return_5d": round((closes[-1] / closes[-6] - 1) * 100, 2),
|
||
"return_10d": round((closes[-1] / closes[-11] - 1) * 100, 2),
|
||
"above_ma20": int(closes[-1] > statistics.fmean(closes[-20:])),
|
||
"volume_ratio_5d": round(volumes[-1] / previous_volume, 2) if previous_volume else 0,
|
||
"volatility_10d": round(statistics.pstdev(returns_10), 2),
|
||
"amount_billion": round(
|
||
_number(current["amount"]) / (100000000 if realtime else 100000), 2
|
||
),
|
||
"turnover_rate": round(
|
||
_number(realtime.get("turnover_rate"))
|
||
if realtime else _number(indicator.get("turnover_rate")),
|
||
2,
|
||
),
|
||
"circ_mv_billion": round(_number(indicator.get("circ_mv")) / 10000, 2),
|
||
"net_flow_million": round(_number(flow.get("net_mf_amount")) / 100, 2),
|
||
"large_flow_million": round(_number(flow.get("large_net_amount")) / 100, 2),
|
||
"limit_status": status,
|
||
"limit_streak": streak,
|
||
}
|
||
)
|
||
|
||
market_return = statistics.fmean(row["return_5d"] for row in factors) if factors else 0
|
||
sectors: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||
for row in factors:
|
||
sectors[row["sector"]].append(row)
|
||
for sector_rows in sectors.values():
|
||
average_return = statistics.fmean(row["return_5d"] for row in sector_rows)
|
||
limit_count = sum(row["limit_status"] == "涨停" or row["pct_chg"] >= 9.5 for row in sector_rows)
|
||
up_count = sum(row["pct_chg"] >= 5 for row in sector_rows)
|
||
strength = min(100, max(0, 50 + average_return * 4 + limit_count * 3 + up_count * 0.6))
|
||
for row in sector_rows:
|
||
row["sector_strength"] = round(strength, 1)
|
||
row["sector_limit_count"] = limit_count
|
||
row["sector_up_count"] = up_count
|
||
row["relative_strength"] = round(row["return_5d"] - market_return, 2)
|
||
return factors, actual_date
|
||
|
||
def apply_formula(
|
||
self, rows: list[dict[str, Any]], formula: dict[str, Any], regime: str
|
||
) -> list[dict[str, Any]]:
|
||
universe = formula["universe"]
|
||
eligible = []
|
||
for row in rows:
|
||
name = str(row.get("name") or "")
|
||
if universe.get("exclude_st") and ("ST" in name.upper() or "退" in name):
|
||
continue
|
||
if row.get("listed_days", 0) < universe.get("listed_days_min", 0):
|
||
continue
|
||
if all(_matches(row.get(item["field"], 0), item["op"], item["value"]) for item in formula["filters"]):
|
||
eligible.append(row)
|
||
if not eligible:
|
||
return []
|
||
|
||
percentiles = {
|
||
item["field"]: _percentile_map(eligible, item["field"], item["direction"])
|
||
for item in formula["score"]
|
||
}
|
||
weight_total = sum(item["weight"] for item in formula["score"])
|
||
results = []
|
||
for row in eligible:
|
||
contributions = []
|
||
score = 0.0
|
||
for item in formula["score"]:
|
||
percentile = percentiles[item["field"]].get(row["ts_code"], 0.5)
|
||
points = percentile * item["weight"] / weight_total
|
||
score += points
|
||
contributions.append(
|
||
{
|
||
"field": item["field"],
|
||
"label": FACTOR_FIELDS[item["field"]],
|
||
"value": row.get(item["field"], 0),
|
||
"points": round(points * 100, 1),
|
||
}
|
||
)
|
||
if score < formula["min_score"]:
|
||
continue
|
||
contributions.sort(key=lambda item: item["points"], reverse=True)
|
||
item = dict(row)
|
||
item["score"] = round(score, 4)
|
||
item["score_display"] = round(score * 100, 1)
|
||
item["contributions"] = contributions
|
||
item["reason"] = "、".join(entry["label"] for entry in contributions[:3])
|
||
item["risk_flags"] = _risk_flags(row, regime)
|
||
results.append(item)
|
||
results.sort(key=lambda item: item["score"], reverse=True)
|
||
return results[: formula["limit"]]
|
||
|
||
def backtest(self, trade_date: str, formula: dict[str, Any]) -> dict[str, Any]:
|
||
dates = self.database.factor_dates(trade_date, 55)
|
||
evaluation_dates = dates[20:-3][-8:]
|
||
wins = 0
|
||
losses = 0
|
||
samples = 0
|
||
returns = []
|
||
drawdowns = []
|
||
all_data = self.database.load_factor_data(trade_date, 60)
|
||
bars_by_code: dict[str, list[dict[str, Any]]] = defaultdict(list)
|
||
for row in all_data["bars"]:
|
||
bars_by_code[row["ts_code"]].append(row)
|
||
for bars in bars_by_code.values():
|
||
bars.sort(key=lambda item: item["trade_date"])
|
||
|
||
for current_date in evaluation_dates:
|
||
try:
|
||
factors, _ = self.build_factors(current_date)
|
||
except ValueError:
|
||
continue
|
||
selected = self.apply_formula(factors, {**formula, "limit": min(10, formula["limit"])}, "backtest")
|
||
for candidate in selected:
|
||
bars = bars_by_code.get(candidate["ts_code"], [])
|
||
index = next((i for i, row in enumerate(bars) if row["trade_date"] == current_date), -1)
|
||
future = bars[index + 1:index + 4] if index >= 0 else []
|
||
if len(future) < 3:
|
||
continue
|
||
entry = candidate["price"]
|
||
won = False
|
||
lost = False
|
||
for day in future:
|
||
low_return = (_number(day["low"]) / entry - 1) * 100
|
||
high_return = (_number(day["high"]) / entry - 1) * 100
|
||
if low_return <= -3:
|
||
lost = True
|
||
break
|
||
if high_return >= 3:
|
||
won = True
|
||
break
|
||
if won:
|
||
wins += 1
|
||
elif lost:
|
||
losses += 1
|
||
samples += 1
|
||
returns.append((_number(future[-1]["close"]) / entry - 1) * 100)
|
||
drawdowns.append(min((_number(day["low"]) / entry - 1) * 100 for day in future))
|
||
return {
|
||
"samples": samples,
|
||
"wins": wins,
|
||
"losses": losses,
|
||
"win_rate": round(wins / samples * 100, 1) if samples else 0,
|
||
"average_3d_return": round(statistics.fmean(returns), 2) if returns else 0,
|
||
"average_drawdown": round(statistics.fmean(drawdowns), 2) if drawdowns else 0,
|
||
"evaluation_days": len(evaluation_dates),
|
||
"definition": "收盘后选股,未来3日先触及+3%且未先触及-3%计为成功;同日双触发按失败处理。",
|
||
"approximate": True,
|
||
}
|
||
|
||
|
||
def compile_local_strategy(prompt: str, regime: str) -> dict[str, Any]:
|
||
base = next((item for item in BUILTIN_STRATEGIES if regime in item["regimes"]), BUILTIN_STRATEGIES[1])
|
||
formula = copy.deepcopy(base["formula"])
|
||
description = prompt.strip() or base["description"]
|
||
lowered = description.lower()
|
||
if "低吸" in description:
|
||
formula["filters"] = [item for item in formula["filters"] if item["field"] != "pct_chg"]
|
||
formula["filters"].append({"field": "pct_chg", "op": "between", "value": [-3, 3]})
|
||
if "放量" in description:
|
||
formula["filters"].append({"field": "volume_ratio_5d", "op": ">=", "value": 1.2})
|
||
if "强势" in description or "突破" in description:
|
||
formula["filters"].append({"field": "return_5d", "op": ">=", "value": 5})
|
||
if "低波" in description or "稳健" in description:
|
||
formula["score"].append({"field": "volatility_10d", "weight": 0.18, "direction": "asc"})
|
||
if "资金" in description or "主力" in description:
|
||
formula["score"].append({"field": "net_flow_million", "weight": 0.18, "direction": "desc"})
|
||
if "小市值" in description or "小盘" in description:
|
||
formula["score"].append({"field": "circ_mv_billion", "weight": 0.15, "direction": "asc"})
|
||
if "少量" in description or "精选" in description:
|
||
formula["limit"] = min(formula["limit"], 8)
|
||
formula["score"] = formula["score"][:12]
|
||
return {
|
||
"name": f"{REGIMES.get(regime, regime)}自定义策略",
|
||
"description": description,
|
||
"regimes": [regime],
|
||
"formula": formula,
|
||
"compiler": "local_template",
|
||
}
|
||
|
||
|
||
def _matches(actual: Any, operator: str, expected: Any) -> bool:
|
||
try:
|
||
if operator == "between":
|
||
return float(expected[0]) <= float(actual) <= float(expected[1])
|
||
if operator == "in":
|
||
return actual in expected
|
||
if operator == ">":
|
||
return float(actual) > float(expected)
|
||
if operator == ">=":
|
||
return float(actual) >= float(expected)
|
||
if operator == "<":
|
||
return float(actual) < float(expected)
|
||
if operator == "<=":
|
||
return float(actual) <= float(expected)
|
||
if operator == "==":
|
||
return actual == expected or float(actual) == float(expected)
|
||
if operator == "!=":
|
||
return actual != expected
|
||
except (TypeError, ValueError, IndexError):
|
||
return False
|
||
return False
|
||
|
||
|
||
def _percentile_map(rows: list[dict[str, Any]], field: str, direction: str) -> dict[str, float]:
|
||
ordered = sorted(rows, key=lambda item: _number(item.get(field)))
|
||
denominator = max(1, len(ordered) - 1)
|
||
result = {}
|
||
for index, row in enumerate(ordered):
|
||
percentile = index / denominator
|
||
result[row["ts_code"]] = 1 - percentile if direction == "asc" else percentile
|
||
return result
|
||
|
||
|
||
def _risk_flags(row: dict[str, Any], regime: str) -> list[str]:
|
||
flags = []
|
||
if row.get("pct_chg", 0) >= 9.5:
|
||
flags.append("当日接近涨停,次日存在高开与无法成交风险")
|
||
if row.get("return_10d", 0) >= 25:
|
||
flags.append("短期累计涨幅较高")
|
||
if row.get("volatility_10d", 0) >= 7:
|
||
flags.append("波动率偏高")
|
||
if row.get("amount_billion", 0) < 1:
|
||
flags.append("成交承载力偏弱")
|
||
if regime == "retreat":
|
||
flags.append("市场处于退潮阶段,策略可能选择空仓")
|
||
return flags
|
||
|
||
|
||
def _regime_reason(regime: str) -> str:
|
||
return {
|
||
"ice": "情绪和赚钱效应处于低位,重点观察率先抗跌与转折信号。",
|
||
"repair": "核心指标从低位改善,适合观察率先修复且有板块共振的方向。",
|
||
"fermentation": "赚钱效应扩散,主线和梯队持续增强。",
|
||
"climax": "情绪处于高位,后排跟风与兑现风险同时上升。",
|
||
"divergence": "指数或核心仍强,但广度、封板质量开始分化。",
|
||
"retreat": "情绪指标继续走弱,应提高筛选门槛并接受无候选结果。",
|
||
}.get(regime, "市场阶段待确认。")
|
||
|
||
|
||
def _number(value: Any, default: float = 0.0) -> float:
|
||
try:
|
||
number = float(value)
|
||
return number if math.isfinite(number) else default
|
||
except (TypeError, ValueError):
|
||
return default
|
||
|
||
|
||
def _display_date(value: str) -> str:
|
||
return f"{value[:4]}-{value[4:6]}-{value[6:8]}" if len(value) == 8 else value
|