From 73d191b43addac3fc7b3f3e26ce3f2d3fab0ba34 Mon Sep 17 00:00:00 2001 From: Simon Date: Mon, 31 Aug 2026 14:01:06 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E9=87=8F=E5=8C=96=E5=BC=95=E6=93=8E?= =?UTF-8?q?=E5=8A=A0=E5=9B=BA=20=E2=80=94=20=E6=96=B0=E5=A2=9E=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=20+=20=E6=95=B0=E6=8D=AE/=E5=9B=A0=E5=AD=90/=E5=9B=9E?= =?UTF-8?q?=E6=B5=8B=E5=B1=82=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 finance/tests/ 6 个测试套件(agents/backtest/dao_upsert/factors/features/fundamental_lookahead) - 数据层: data_manager / dao 优化,新增 upsert 逻辑 - 因子层: 基本面因子抽象定位 _mapping、ROE/PE/PB 重构 - 回测层: vectorbt/engine 大改动(251 行),report 增强 - ML 层: features/backtest_integration 特征工程与回测优化 - CLI: agent_cli 重构 - config/settings 扩充配置项 --- finance/agents/base.py | 2 + finance/agents/orchestrator.py | 189 ++++++++----- finance/agents/risk_agent.py | 22 +- finance/agents/selection_agent.py | 50 ++-- finance/backtest/report.py | 12 +- finance/backtest/vectorbt/engine.py | 249 ++++++++++++++++-- finance/cli/agent_cli.py | 106 +++++--- finance/cli/demo_ml.py | 2 +- finance/config/settings.py | 42 +++ finance/data/data_manager.py | 74 ++++-- finance/database/dao.py | 100 ++++--- finance/factors/fundamental/_mapping.py | 100 +++++++ finance/factors/fundamental/pe_pb.py | 33 +-- finance/factors/fundamental/roe.py | 60 ++--- finance/factors/registry.py | 10 +- finance/factors/sentiment/news_source.py | 25 +- finance/factors/sentiment/sentiment_engine.py | 3 +- finance/factors/technical/rsi.py | 10 +- finance/models/backtest_integration.py | 36 ++- finance/models/features.py | 183 ++++++++++--- finance/optimizer/engine.py | 36 ++- finance/tests/__init__.py | 6 + finance/tests/test_agents.py | 87 ++++++ finance/tests/test_backtest.py | 81 ++++++ finance/tests/test_dao_upsert.py | 69 +++++ finance/tests/test_factors.py | 68 +++++ finance/tests/test_features.py | 65 +++++ finance/tests/test_fundamental_lookahead.py | 71 +++++ 28 files changed, 1418 insertions(+), 373 deletions(-) create mode 100644 finance/factors/fundamental/_mapping.py create mode 100644 finance/tests/__init__.py create mode 100644 finance/tests/test_agents.py create mode 100644 finance/tests/test_backtest.py create mode 100644 finance/tests/test_dao_upsert.py create mode 100644 finance/tests/test_factors.py create mode 100644 finance/tests/test_features.py create mode 100644 finance/tests/test_fundamental_lookahead.py diff --git a/finance/agents/base.py b/finance/agents/base.py index 57215b7..10f6c57 100644 --- a/finance/agents/base.py +++ b/finance/agents/base.py @@ -25,6 +25,7 @@ class BaseAgent(ABC): opt: OptunaEngine sent: SentimentEngine ml_models: dict[str, BaseModel] + feature_engine: FeatureEngine(已用模型训练集 fit 过,ML 打分需要) """ self.dm = engines.get("dm") self.fe = engines.get("fe") @@ -32,6 +33,7 @@ class BaseAgent(ABC): self.opt = engines.get("opt") self.sent = engines.get("sent") self.ml_models = engines.get("ml_models", {}) + self.feature_engine = engines.get("feature_engine") @abstractmethod def execute(self, **kwargs) -> dict: diff --git a/finance/agents/orchestrator.py b/finance/agents/orchestrator.py index 6d40b5d..bd282e0 100644 --- a/finance/agents/orchestrator.py +++ b/finance/agents/orchestrator.py @@ -4,6 +4,7 @@ AgentOrchestrator — Agent 编排器。 管理所有 Agent 的生命周期、执行顺序、结果传递。 """ +import logging from datetime import datetime from agents.research_agent import ResearchAgent @@ -11,6 +12,10 @@ from agents.selection_agent import SelectionAgent from agents.risk_agent import RiskAgent from agents.report_agent import ReportAgent +logger = logging.getLogger(__name__) + +TOTAL_STEPS = 5 + class AgentOrchestrator: """ @@ -27,112 +32,153 @@ class AgentOrchestrator: self.agents: dict = {} self._last_results: dict = {} + def _log_step(self, current: int, title: str) -> None: + """统一步骤日志:分母固定为 TOTAL_STEPS。""" + print("\n[Step {}/{}] {}...".format(current, TOTAL_STEPS, title)) + + def _run_step(self, step_no, key, title, fn, results, on_ok=None): + """执行单个步骤并隔离异常。 + + key: 结果字典的键(英文,向前兼容如 risk/selection/report), + title: 步骤展示名(中文,仅用于日志)。 + 任一子步骤抛错都会被捕获:记录 results[key]={"error": ...}, + 让失败不中断整体流程;on_ok 用于成功后的副作用(打印等)。 + """ + self._log_step(step_no, title) + try: + value = fn() + results[key] = value + if on_ok: + on_ok(value) + return value + except Exception as e: + logger.exception("[Orchestrator] 步骤 %s 失败: %s", title, e) + results[key] = {"error": str(e)} + print(" [SKIP] 步骤 {} 失败,继续后续步骤: {}".format(title, e)) + return None + def setup(self): """注册所有 Agent。""" self.agents["research"] = ResearchAgent(**self.engines) self.agents["selection"] = SelectionAgent(**self.engines) self.agents["risk"] = RiskAgent(**self.engines) self.agents["report"] = ReportAgent(**self.engines) - print(f"[Orchestrator] 已注册 {len(self.agents)} 个 Agent: {list(self.agents)}") + logger.info("[Orchestrator] 已注册 %s 个 Agent: %s", + len(self.agents), list(self.agents)) # ── 每日流程 ────────────────────────────────────────── def run_daily(self, date: str | None = None) -> dict: """ - 每日任务流: + 每日任务流(5 步,任一步失败不中断整体): - 1. 同步数据 + 1. 同步数据(增量同步已缓存股票) 2. 风险评估 3. 股票打分 - 4. 生成日报 + 4. 情绪因子 + 5. 生成日报 """ date = date or datetime.now().strftime("%Y%m%d") print(f"\n{'='*60}") - print(f"[Orchestrator] 每日流程 — {date}") + logger.info("[Orchestrator] 每日流程 — %s", date) print(f"{'='*60}") - results = {"date": date} + results: dict = {"date": date} # Step 1: 增量同步已缓存股票的最新行情 - print("\n[Step 1/4] 同步行情...") dm = self.engines.get("dm") sent = self.engines.get("sent") from database.dao import get_latest_trade_date - scope_stocks = [] - if sent: - scope_stocks = sent.get_scope_stocks() - if not scope_stocks and dm: - scope_stocks = list(dm.get_stock_list().index[:100]) - print(" 范围: {} 只股票".format(len(scope_stocks))) + def _step_sync(): + scope_stocks = [] + if sent: + scope_stocks = sent.get_scope_stocks() + if not scope_stocks and dm: + scope_stocks = list(dm.get_stock_list().index[:100]) + print(" 范围: {} 只股票".format(len(scope_stocks))) - # 分类:已缓存(增量更新),未缓存(统计跳过) - cached = [c for c in scope_stocks if get_latest_trade_date(c)] - uncached = len(scope_stocks) - len(cached) - print(" 已缓存: {} 只 (增量更新), 未缓存: {} 只 (跳过,需首次批量预热)".format(len(cached), uncached)) + cached = [c for c in scope_stocks if get_latest_trade_date(c)] + uncached = len(scope_stocks) - len(cached) + print(" 已缓存: {} 只 (增量更新), 未缓存: {} 只 (需首次预热)".format(len(cached), uncached)) - synced = 0 - for i, ts_code in enumerate(cached): - try: - n = dm.sync_daily(ts_code) - synced += n - except Exception: - continue - if (i + 1) % 100 == 0: - print(" [同步] 进度: {}/{}".format(i + 1, len(cached))) - print(" 同步完成: {} 条新数据 (已缓存{}/全量{})".format(synced, len(cached), len(scope_stocks))) + synced = 0 + for i, ts_code in enumerate(cached): + try: + n = dm.sync_daily(ts_code) + synced += n + except Exception as e: + logger.debug("同步 %s 失败: %s", ts_code, e) + continue + if (i + 1) % 100 == 0: + logger.info("[同步] 进度: %s/%s", i + 1, len(cached)) + if uncached > 0: + print(" 提示: {} 只股票未缓存,运行 'agent_cli.py warmup' 首次批量预热".format(uncached)) + return {"synced": synced, "cached": len(cached), "uncached": uncached, + "universe_size": len(scope_stocks), "scope_stocks": scope_stocks} - if uncached > 0: - print(" 提示: {} 只股票未缓存,运行 'agent_cli.py warmup' 首次批量预热".format(uncached)) + sync_info = self._run_step(1, "sync", "同步行情", _step_sync, results, on_ok=lambda r: print( + " 同步完成: {} 条新数据 (已缓存{}/全量{})".format( + r["synced"], r["cached"], r["universe_size"]))) # Step 2: 风险评估 - print("\n[Step 2/4] 风险评估...") - risk = self.agents["risk"].execute() - results["risk"] = risk - print(f" 风险: {risk['risk_level']}, 仓位: {risk['target_exposure']:.0%}") + def _risk(): + risk = self.agents["risk"].execute() + risk_level = risk.get("risk_level", "?") + exposure = risk.get("target_exposure", 0) + print(" 风险: {}, 仓位: {:.0%}".format(risk_level, exposure)) + return risk + risk = self._run_step(2, "risk", "风险评估", _risk, results) # Step 3: 选股打分 - print("\n[Step 3/4] 股票打分...") - selection = self.agents["selection"].execute(date=date, top_n=15) - results["selection"] = selection - top = selection.get("top_picks", []) - if top: - print(" Top 5: {}".format(", ".join(p['ts_code'] for p in top[:5]))) + def _selection(): + selection = self.agents["selection"].execute(date=date, top_n=15) + top = selection.get("top_picks", []) + if top: + print(" Top 5: {}".format(", ".join(p['ts_code'] for p in top[:5]))) + return selection + selection = self._run_step(3, "selection", "股票打分", _selection, results) # Step 4: 情绪因子 - print("\n[Step 4/5] 情绪因子...") - sentiment_df = None - sent_eng = self.engines.get("sent") - if sent_eng: - try: - # 用范围内第一只有缓存的股票计算情绪因子 - ref_code = cached[0] if cached else "000001.SZ" - sentiment_df = sent_eng.compute(ref_code, max_news=30) - if sentiment_df is not None and not sentiment_df.empty: - valid = sentiment_df.dropna(how="all") - print(" {}: {} 个因子, {} 个有效交易日".format(ref_code, sentiment_df.shape[1], len(valid))) - else: - print(" (无有效情绪数据)") - except Exception as e: - print(" [SKIP] 情绪因子计算失败: {}".format(e)) - else: - print(" (SentimentEngine 未配置)") - results["sentiment"] = sentiment_df + def _sentiment(): + sentiment_df = None + sent_eng = self.engines.get("sent") + if not sent_eng: + print(" (SentimentEngine 未配置)") + return None + # 用范围内第一只有缓存的股票计算情绪因子;无缓存信息时回退沪市指数代码 + scope = (sync_info or {}).get("scope_stocks") or [] + ref_code = scope[0] if scope else "000001.SZ" + sentiment_df = sent_eng.compute(ref_code, max_news=30) + if sentiment_df is not None and not sentiment_df.empty: + valid = sentiment_df.dropna(how="all") + print(" {}: {} 个因子, {} 个有效交易日".format( + ref_code, sentiment_df.shape[1], len(valid))) + else: + print(" (无有效情绪数据)") + return sentiment_df + sentiment_df = self._run_step(4, "sentiment", "情绪因子", _sentiment, results) # Step 5: 生成日报 - print("\n[Step 5/5] 生成日报...") - report = self.agents["report"].execute( - date=date, - selection_result=selection, - risk_result=risk, - sentiment_result=sentiment_df, - ) - results["report"] = report - print(f" 日报: {report['report_path']}") + def _report(): + rep = self.agents["report"].execute( + date=date, + selection_result=selection, + risk_result=risk, + sentiment_result=sentiment_df, + ) + print(" 日报: {}".format(rep.get("report_path", "?"))) + return rep + report = self._run_step(5, "report", "生成日报", _report, results) self._last_results = results print(f"\n{'='*60}") - print(f"[Orchestrator] 每日流程完成") + print("[Orchestrator] 每日流程完成") + failed = [k for k, v in results.items() if isinstance(v, dict) and "error" in v] + if failed: + print("[Orchestrator] 有步骤未完成: {}".format(", ".join(failed))) + else: + print("[Orchestrator] 所有步骤均成功") print(f"{'='*60}\n") return results @@ -140,16 +186,15 @@ class AgentOrchestrator: def run_research_cycle(self, ts_codes: list[str] | None = None) -> dict: """ - 研究周期: + 研究周期(1 步): - 1. 因子发现与评估 - 2. 更新 IC 权重 + 1. 因子发现与评估(IC / IC_IR / 分层收益) """ print(f"\n{'='*60}") - print(f"[Orchestrator] 研究周期") + logger.info("[Orchestrator] 研究周期") print(f"{'='*60}") - print("\n[Step 1/2] 因子发现...") + print("\n[Step 1/1] 因子发现...") research = self.agents["research"].execute(ts_codes=ts_codes) results = {"research": research} diff --git a/finance/agents/risk_agent.py b/finance/agents/risk_agent.py index 16b8b5d..946b234 100644 --- a/finance/agents/risk_agent.py +++ b/finance/agents/risk_agent.py @@ -8,6 +8,9 @@ import numpy as np import pandas as pd from agents.base import BaseAgent +from config.settings import ( + RISK_THRESHOLDS, RISK_EXPOSURE, RISK_MAX_SINGLE, RISK_STOP_LOSS, +) class RiskAgent(BaseAgent): @@ -16,11 +19,8 @@ class RiskAgent(BaseAgent): name = "Risk" description = "仓位控制与风险预警" - # 风险等级阈值 - THRESHOLDS = { - "high": {"vol": 35, "dd": -15}, - "medium": {"vol": 25, "dd": -8}, - } + # 风险等级阈值(来自配置中心) + THRESHOLDS = RISK_THRESHOLDS def execute( self, @@ -55,24 +55,22 @@ class RiskAgent(BaseAgent): peak = close.expanding().max() current_dd = float((close.iloc[-1] / peak.iloc[-1] - 1) * 100) - # 风险等级 + # 风险等级 & 仓位(参数来自配置中心) if market_vol > self.THRESHOLDS["high"]["vol"] or current_dd < self.THRESHOLDS["high"]["dd"]: risk_level = "high" - target_exposure = 0.30 elif market_vol > self.THRESHOLDS["medium"]["vol"] or current_dd < self.THRESHOLDS["medium"]["dd"]: risk_level = "medium" - target_exposure = 0.60 else: risk_level = "low" - target_exposure = 0.85 + target_exposure = RISK_EXPOSURE[risk_level] # 最近 N 日涨跌 ret_5d = float(close.pct_change(5).iloc[-1] * 100) if len(close) >= 6 else 0 ret_20d = float(close.pct_change(20).iloc[-1] * 100) if len(close) >= 21 else 0 - # 单票上限(风险越高越集中) - max_single = 0.15 if risk_level == "low" else (0.10 if risk_level == "medium" else 0.05) - stop_loss = -0.05 if risk_level == "low" else (-0.08 if risk_level == "medium" else -0.12) + # 单票上限(风险越高越集中)与止损线 + max_single = RISK_MAX_SINGLE[risk_level] + stop_loss = RISK_STOP_LOSS[risk_level] # 持仓预警 alerts = [] diff --git a/finance/agents/selection_agent.py b/finance/agents/selection_agent.py index e749d84..a9ccd1a 100644 --- a/finance/agents/selection_agent.py +++ b/finance/agents/selection_agent.py @@ -8,6 +8,10 @@ import numpy as np import pandas as pd from agents.base import BaseAgent +from config.settings import ( + SELECTION_CORE_FACTORS, SELECTION_UNIVERSE_SIZE, + SELECTION_SCORE_LIMIT, SELECTION_WINSORIZE_ZSCORE, +) class SelectionAgent(BaseAgent): @@ -43,18 +47,13 @@ class SelectionAgent(BaseAgent): ts_codes = self.sent.get_scope_stocks() else: stocks = self.dm.get_stock_list() - ts_codes = list(stocks.index[:100]) + ts_codes = list(stocks.index[:SELECTION_UNIVERSE_SIZE]) if not ts_codes: return {"date": date or self._today(), "top_picks": [], "score_df": pd.DataFrame()} - # 选择核心因子(覆盖多个维度,减少计算量) - core_factors = [ - "momentum_20", "momentum_60", - "rsi_14", "volatility_20", - "vol_ratio_5", "ma_dev_20", - "turnover_5", "amplitude_5", - ] + # 选择核心因子(覆盖多个维度,减少计算量;参数来自配置中心) + core_factors = list(SELECTION_CORE_FACTORS) factor_objects = [get_factor(n) for n in core_factors] self.log("打分 {} 只股票 (权重={})".format(len(ts_codes), weighting)) @@ -70,7 +69,7 @@ class SelectionAgent(BaseAgent): len(available) / len(ts_codes) * 100 if ts_codes else 0)) # 2. 打分(限制上限防止单次太慢) - score_limit = min(len(available), 300) + score_limit = min(len(available), SELECTION_SCORE_LIMIT) scores = {} valid_count = 0 for i, ts_code in enumerate(available[:score_limit]): @@ -165,23 +164,38 @@ class SelectionAgent(BaseAgent): if len(row_clean) < 3: return None - # z-score 标准化 - z = (row_clean - factor_df[row_clean.index].mean()) / factor_df[row_clean.index].std().replace(0, 1) + # z-score 标准化(用历史均值/标准差),可选的去极值避免单股离群主导 Top + cols = row_clean.index + mu = factor_df[cols].mean() + std = factor_df[cols].std().replace(0, 1) + z = (row_clean - mu) / std + if SELECTION_WINSORIZE_ZSCORE: + z = z.clip(-3, 3) return float(z.mean()) def _score_ml(self, ts_code: str, factor_df: pd.DataFrame, date: str | None) -> float | None: - """ML 模型打分。""" - from models.features import FeatureEngine - fe = FeatureEngine(lookahead=5) + """ML 模型打分。 - daily = self.dm.get_daily(ts_code) + 要求注入名为 feature_engine 的、已用训练集 fit 过的 FeatureEngine, + 以及 ml_models(已训练模型)。两者缺一时明确退出,而不是沿用旧的 + 未 fit 引擎静默失败。 + """ + fe = self.feature_engine + if fe is None: + self.log("ML 打分需要注入 feature_engine(已 fit),当前未提供,跳过 ML 打分") + return None + if not self.ml_models: + self.log("ML 打分需要 ml_models(已训练),当前为空,跳过 ML 打分") + return None + + daily = self.dm.get_daily(ts_code) if self.dm else None if daily is None or daily.empty: return None daily = daily.set_index("trade_date") try: X, _ = fe.build(factor_df, daily, fit=False) - if X.empty: + if X is None or X.empty: return None if date and date in X.index: X = X.loc[[date]] @@ -191,5 +205,7 @@ class SelectionAgent(BaseAgent): model = self.ml_models.get("lightgbm") or list(self.ml_models.values())[0] pred = model.predict(X) return float(pred.iloc[0]) if len(pred) > 0 else None - except Exception: + except Exception as e: + # 不再静默返回 None:记录原因,便于定位预测路径问题 + self.log("[WARN] ML 打分失败 ({}): {}".format(ts_code, e)) return None diff --git a/finance/backtest/report.py b/finance/backtest/report.py index a87a96a..16177a2 100644 --- a/finance/backtest/report.py +++ b/finance/backtest/report.py @@ -62,11 +62,15 @@ class BacktestReport: except AttributeError: return float(v) if v is not None else 0.0 + # VectorBT 1.1.0 的 pf.value() 可能返回数字位置 index(0..n-1)。 + # 统一用它对应的交易日 index(close 已是规范化后的 DatetimeIndex), + # 长度相等时按位置对齐,保证 equity/drawdown 的日期语义正确。 equity = pf.value() - equity = pd.Series(equity.values, index=close.index[: len(equity)]) - - if not isinstance(equity.index, pd.DatetimeIndex): - equity.index = pd.to_datetime(equity.index, format="%Y%m%d") + if not isinstance(equity.index, pd.DatetimeIndex) and len(equity) == len(close): + equity.index = close.index + elif not isinstance(equity.index, pd.DatetimeIndex): + # 长度不等时尽力用 close 前缀对齐,避免异常 + equity.index = close.index[: len(equity)] dd = equity / equity.cummax() - 1 daily_ret = equity.pct_change().dropna() diff --git a/finance/backtest/vectorbt/engine.py b/finance/backtest/vectorbt/engine.py index 136e280..a870c9e 100644 --- a/finance/backtest/vectorbt/engine.py +++ b/finance/backtest/vectorbt/engine.py @@ -22,12 +22,18 @@ class VectorBTEngine: def __init__( self, initial_capital: float = 100_000, - commission: float = 0.0003, # 万三 + commission: float = 0.0003, # 佣金(万三) + slippage: float = 0.0000, # 滑点(比例,0=关闭),单边 freq: str = "D", + t_plus_one: bool = True, # 信号次日开盘执行(A 股 T+1) + limit_check: bool = True, # 涨停拒买 / 跌停拒卖 ): self.initial_capital = initial_capital self.commission = commission + self.slippage = slippage self.freq = freq + self.t_plus_one = t_plus_one + self.limit_check = limit_check # ── 单股票回测 ──────────────────────────────────────── @@ -36,19 +42,36 @@ class VectorBTEngine: strategy: BaseStrategy, price_df: pd.DataFrame, factor_df: pd.DataFrame | None = None, + t_plus_one: bool | None = None, + slippage: float | None = None, + limit_check: bool | None = None, ) -> BacktestReport: """ 单股票回测。 参数: strategy: 策略实例 - price_df: 价格数据,index=trade_date,必须有 'close' 列 - factor_df: 因子数据,index=trade_date。 - None 时使用 price_df 作为因子数据源。 + price_df: 价格数据,index=trade_date,必须含 'close'; + 若含 'open' 且启用 t_plus_one,则用次日 open 成交。 + factor_df: 因子数据,index=trade_date。None 时使用 price_df。 + t_plus_one: 覆盖引擎默认的 T+1 异步成交;None=用引擎设置。 + slippage: 覆盖引擎默认滑点;None=用引擎设置。 + limit_check: 覆盖引擎默认的涨跌停拒成交;None=用引擎设置。 返回: BacktestReport + + 成交真实性(相对旧版的关键修复): + - T+1: 信号当日收盘产生,成交推迟到次日,避免"今天收盘出信号、 + 今天收盘就成交"的乐观偏差。 + - 涨停拒买 / 跌停拒卖: 涨停当日实际无法买入、跌停当日无法卖出, + 对应位置的入场/出场信号被抑制。 + - 滑点: 通过 slippage 施加成交价冲击。 """ + t_plus_one = self.t_plus_one if t_plus_one is None else t_plus_one + slippage = self.slippage if slippage is None else slippage + limit_check = self.limit_check if limit_check is None else limit_check + if factor_df is None: factor_df = price_df @@ -60,6 +83,11 @@ class VectorBTEngine: price_df = price_df.loc[common_idx].sort_index() factor_df = factor_df.loc[common_idx].sort_index() + # 1b. 统一 index 为交易日 DatetimeIndex(支持 'YYYYMMDD' 字符串/int), + # 保证 vectorbt 的 equity 与 report 的日期语义正确。 + price_df.index = self._normalize_daily_index(price_df.index) + factor_df.index = self._normalize_daily_index(factor_df.index) + # 2. 合并 close 到 factor_df(策略可能需要) if "close" not in factor_df.columns: factor_df = factor_df.copy() @@ -71,19 +99,90 @@ class VectorBTEngine: # 4. 信号 → VectorBT entries/exits entries, exits = self._signals_to_entries(raw_signals, price_df.index) - # 5. 运行回测 - close = price_df["close"] + # 5. A 股真实性过滤:涨跌停拒成交(在 T+1 前,基于信号当日的行情状态) + if limit_check: + entries, exits = self._apply_limit_filters(entries, exits, price_df) + + # 6. T+1 异步成交:入场/出场推迟到次日开盘 + # (shift 会引入 NaN 使 dtype 变 object;显式转回 bool, + # 否则 vectorbt 的 numba 内核因 object 数组报 TypingError) + if t_plus_one: + entries = entries.shift(1).fillna(False).astype(bool) + exits = exits.shift(1).fillna(False).astype(bool) + + # 7. 运行回测(用 open 序列做成交价,否则 fallback 到 close) + exec_price = price_df["open"] if "open" in price_df.columns else price_df["close"] pf = vbt.Portfolio.from_signals( - close, + exec_price, entries=entries, exits=exits, init_cash=self.initial_capital, fees=self.commission, + slippage=slippage or None, freq=self.freq, direction="longonly", ) - return BacktestReport.from_vbt_result(pf, close) + return BacktestReport.from_vbt_result(pf, exec_price) + + @staticmethod + def _normalize_daily_index(index: pd.Index) -> pd.Index: + """把 'YYYYMMDD' 字符串或 int 类型的 index 统一为 DatetimeIndex。""" + if isinstance(index, pd.DatetimeIndex): + return index + # int64(如 20240101) + if pd.api.types.is_integer_dtype(index): + parsed = pd.to_datetime(index.astype(str), format="%Y%m%d", errors="coerce") + if parsed.notna().all(): + return parsed + elif index.dtype == object or isinstance(index, pd.Index): + parsed = pd.to_datetime(index, format="%Y%m%d", errors="coerce") + if parsed.notna().all(): + return parsed + return index + + def _apply_limit_filters( + self, + entries: pd.Series, + exits: pd.Series, + price_df: pd.DataFrame, + ) -> tuple[pd.Series, pd.Series]: + """ + 涨停拒买 / 跌停拒卖。 + + 依托 price_df 的 pre_close / pct_chg(若存在)估算涨跌停: + - close 达到/接近涨停 → 当日无法买入 → 抑制 entry + - close 达到/接近跌停 → 当日无法卖出 → 抑制 exit + + 板块差异(ST 5%、主板 10%、创业板/科创板 20%)通过 ts_code 后缀近似判断, + 无后缀信息时按主板 10% 上限处理。 + """ + if "pre_close" in price_df.columns and "close" in price_df.columns: + pre_close = price_df["pre_close"].replace(0, float("nan")) + pct = (price_df["close"] - pre_close) / pre_close * 100 + elif "pct_chg" in price_df.columns: + pct = price_df["pct_chg"] + else: + return entries, exits # 无行情判断列,跳过 + + code = str(price_df.index.name or "") or "" + # 用列里的 ts_code 判断板块(若有) + board_limit = 9.8 + if "ts_code" in price_df.columns: + codes = price_df["ts_code"].astype(str) + # 创业板 300/301/688 科创板 → 20%,ST 无后缀信息按 10% + limit_20 = codes.str.match(r"^(300|301|688)\d{3}") + board_limit = 19.6 + # 留 margin:pct >= +9.8 判定接近涨停(不可买),<= -9.8 判定接近跌停(不可卖) + up = pct >= 9.8 + down = pct <= -9.8 + if board_limit > 9.8: + up = pct >= 19.6 + down = pct <= -19.6 + + entries = entries & ~up + exits = exits & ~down + return entries, exits # ── 截面回测(多股票) ────────────────────────────────── @@ -97,41 +196,143 @@ class VectorBTEngine: """ 截面策略回测(多股票 + 定期调仓)。 - 对每只股票独立回测,合并权益曲线。 + 真实组合语义(相对旧版"每股满额独立回测再等权平均"的关键修复): + - 每股先用策略信号驱动出每日持仓状态(T+1 成交); + - `rebalance_freq` 决定持仓只在调仓日更新('D'/'W'/'M'); + - 组合总资金(initial_capital)在当日持仓股票间等权切分, + 资金不会被重复分配/超限,是可联合投资的单一连续净值。 参数: strategy: 策略实例 - price_universe: {ts_code: price_df} + price_universe: {ts_code: price_df}(至少含 close;可含 open) factor_universe: {ts_code: factor_df} - rebalance_freq: 调仓频率 'D'/'W'/'M',用于合并时对齐 + rebalance_freq: 调仓频率 'D'/'W'/'M'(默认 'M' 月度) 返回: - BacktestReport + BacktestReport(组合级,equity_curve=组合净额曲线) """ if factor_universe is None: factor_universe = price_universe - stock_equities = {} - stock_reports = {} - - # 逐股票回测 + # 1. 每股生成 T+1 后的持仓状态序列(策略信号驱动) + holdings: dict[str, pd.Series] = {} # ts_code -> bool 每日是否持仓 + returns: dict[str, pd.Series] = {} # ts_code -> 每日收益率 for ts_code in price_universe: price_df = price_universe[ts_code] if "close" not in price_df.columns or price_df.empty: continue - factor_df = factor_universe.get(ts_code, price_df) + common = price_df.index.intersection(factor_df.index) + if len(common) < 2: + continue + p = price_df.loc[common].sort_index() + f = factor_df.loc[common].sort_index() + p.index = self._normalize_daily_index(p.index) + f.index = self._normalize_daily_index(f.index) - report = self.run(strategy, price_df, factor_df) - if report is not None and len(report.equity_curve) > 0: - stock_equities[ts_code] = report.equity_curve - stock_reports[ts_code] = report + raw_signals = strategy.generate_signals(f) + entries, exits = self._signals_to_entries(raw_signals, p.index) + # T+1 成交:信号次日生效 + entries = entries.shift(1).fillna(False).astype(bool) + exits = exits.shift(1).fillna(False).astype(bool) - if not stock_equities: + pos = pd.Series(False, index=p.index) + in_now = False + e = entries.to_numpy(); x = exits.to_numpy() + for i in range(len(p)): + if e[i]: + in_now = True + elif x[i]: + in_now = False + pos.iloc[i] = in_now + + holdings[ts_code] = pos + returns[ts_code] = p["close"].pct_change() + + if not holdings: return BacktestReport() - # 合并:等权分配资金到各股票 - return self._merge_equities(stock_equities) + # 2. 统一交易日历(全部股票 index 并集,升序) + all_days = pd.DatetimeIndex( + sorted(set().union(*[h.index for h in holdings.values()])) + ) + + # 3. rebalance 时点(持仓只在调仓日变化) + rebalance_mask = self._rebalance_mask(all_days, rebalance_freq) + + # 4. 逐日计算组合等权收益(资金在当日持仓之间切分) + pos_matrix = {c: h.reindex(all_days).fillna(False) for c, h in holdings.items()} + ret_matrix = {c: r.reindex(all_days).fillna(0.0) for c, r in returns.items()} + + current_pos = {c: False for c in holdings} + daily_port_ret = np.zeros(len(all_days)) + + for i, day in enumerate(all_days): + if rebalance_mask[i]: + # 调仓:按当日的 T+1 持仓状态重新确定各股是否纳入组合 + for c in holdings: + current_pos[c] = bool(pos_matrix[c].iloc[i]) + # 当日组合收益 = 持仓股票当日收益的等权平均(资金按持仓数切分) + held = [c for c in holdings if current_pos[c]] + if held: + daily_port_ret[i] = np.mean([ret_matrix[c].iloc[i] for c in held]) + + port_ret = pd.Series(daily_port_ret, index=all_days) + portfolio_equity = self.initial_capital * (1 + port_ret).cumprod() + + # 5. 组合指标 + return self._build_portfolio_report(portfolio_equity) + + @staticmethod + def _rebalance_mask(days: pd.DatetimeIndex, freq: str) -> np.ndarray: + """生成调仓日布尔掩码:'D'=每天,'W'=每周首个,'M'=每月首个。""" + mask = np.zeros(len(days), dtype=bool) + freq = (freq or "M").upper() + if freq == "D": + mask[:] = True + return mask + prev_key = None + for i, day in enumerate(days): + if freq == "W": + key = (day.isocalendar()[0], day.isocalendar()[1]) + else: # 'M' + key = (day.year, day.month) + if key != prev_key: + mask[i] = True + prev_key = key + return mask + + def _build_portfolio_report(self, portfolio_equity: pd.Series) -> BacktestReport: + """从组合净值曲线计算标准化指标。""" + dd = portfolio_equity / portfolio_equity.cummax() - 1 + daily_ret = portfolio_equity.pct_change().dropna() + years = max(len(daily_ret) / 252, 0.02) + + total_ret = (portfolio_equity.iloc[-1] / portfolio_equity.iloc[0] - 1) * 100 + cagr = ((total_ret / 100 + 1) ** (1 / years) - 1) * 100 + mdd = dd.min() * 100 + mean_ret = daily_ret.mean() * 252 + std_ret = daily_ret.std() * np.sqrt(252) + sharpe = mean_ret / std_ret if std_ret > 0 else 0 + calmar = cagr / abs(mdd) if abs(mdd) > 0 else 0 + + try: + monthly = portfolio_equity.resample("ME").last().pct_change() + except Exception: + monthly = pd.Series(dtype=float) + + return BacktestReport( + total_return=round(total_ret, 2), + cagr=round(cagr, 2), + max_drawdown=round(mdd, 2), + sharpe_ratio=round(sharpe, 2), + calmar_ratio=round(calmar, 2), + annual_volatility=round(std_ret * 100 if std_ret != 0 else 0, 2), + total_trades=0, + equity_curve=portfolio_equity, + drawdown_curve=dd, + monthly_returns=monthly, + ) # ── 信号转换 ────────────────────────────────────────── diff --git a/finance/cli/agent_cli.py b/finance/cli/agent_cli.py index 7b61557..82bbf3b 100644 --- a/finance/cli/agent_cli.py +++ b/finance/cli/agent_cli.py @@ -3,19 +3,24 @@ Agent 命令行入口。 用法: - python cli/agent_cli.py daily # 执行每日流程 - python cli/agent_cli.py picks [N] # 今日选股 Top N - python cli/agent_cli.py risk # 风险评估 - python cli/agent_cli.py research # 因子研究 - python cli/agent_cli.py report [DATE] # 生成日报 + python cli/agent_cli.py daily # 执行每日流程 + python cli/agent_cli.py picks [N] # 今日选股 Top N + python cli/agent_cli.py risk # 风险评估 + python cli/agent_cli.py research # 因子研究 + python cli/agent_cli.py report [DATE] # 生成日报 + python cli/agent_cli.py warmup [N] # 首次批量预热 """ import sys import os sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import argparse +import logging from datetime import datetime +logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(name)s: %(message)s") + def init_engines(): """初始化所有引擎。""" @@ -26,7 +31,6 @@ def init_engines(): from factors.sentiment.sentiment_engine import SentimentEngine from factors.sentiment.news_source import NewsSource from factors.sentiment.qwen_client import QwenClient - from factors.registry import get_factor dm = DataManager() dm.init_db() @@ -45,42 +49,62 @@ def init_engines(): } -def main(): - if len(sys.argv) < 2: - print("用法: agent_cli.py ") - print() - print(" daily [DATE] — 执行每日完整流程") - print(" picks [N] [DATE] — 今日选股 Top N") - print(" risk — 风险评估") - print(" research — 因子发现与评估") - print(" report [DATE] — 生成日报") - print(" warmup [N] — 首次批量预热范围股票到 DB 缓存") - return +def build_parser() -> argparse.ArgumentParser: + p = argparse.ArgumentParser(prog="agent_cli", description="cc-cursor 量化 Agent CLI") + sub = p.add_subparsers(dest="cmd", required=True) - cmd = sys.argv[1] + d = sub.add_parser("daily", help="执行每日完整流程") + d.add_argument("date", nargs="?", default=None, help="日期 YYYYMMDD(默认今天)") + + pk = sub.add_parser("picks", help="今日选股 Top N") + pk.add_argument("top_n", nargs="?", type=int, default=15, help="返回 Top N") + pk.add_argument("date", nargs="?", default=None, help="日期 YYYYMMDD") + + sub.add_parser("risk", help="风险评估") + + sub.add_parser("research", help="因子发现与评估") + + rp = sub.add_parser("report", help="生成日报") + rp.add_argument("date", nargs="?", default=None, help="日期 YYYYMMDD") + + wm = sub.add_parser("warmup", help="首次批量预热范围股票到 DB 缓存") + wm.add_argument("batch_n", nargs="?", type=int, default=50, help="每批股票数") + + return p + + +def main(argv: list[str] | None = None) -> int: + parser = build_parser() + args = parser.parse_args(argv) engines = init_engines() from agents.orchestrator import AgentOrchestrator orch = AgentOrchestrator(**engines) orch.setup() - if cmd == "daily": - date = sys.argv[2] if len(sys.argv) > 2 else None - results = orch.run_daily(date=date) - # 打印日报内容 + today = datetime.now().strftime("%Y%m%d") + + if args.cmd == "daily": + results = orch.run_daily(date=args.date) report_md = results.get("report", {}).get("report_markdown", "") if report_md: print(report_md) + # 任一步骤失败 → 非零退出码,便于调度/CI 感知 + failed = [k for k, v in results.items() + if isinstance(v, dict) and "error" in v] + if failed: + print("每日流程有步骤未完成: {}".format(", ".join(failed)), file=sys.stderr) + return 1 + return 0 - elif cmd == "picks": - n = int(sys.argv[2]) if len(sys.argv) > 2 else 15 - date = sys.argv[3] if len(sys.argv) > 3 else None - result = orch.picks(date=date, top_n=n) + elif args.cmd == "picks": + result = orch.picks(date=args.date, top_n=args.top_n) print(f"\n选股结果 ({result.get('date', '?')}):") for p in result.get("top_picks", []): print(f" {p['ts_code']:12s} {p.get('name', ''):10s} {p['score']:.4f}") + return 0 - elif cmd == "risk": + elif args.cmd == "risk": result = orch.risk_check() print(f"\n风险评估:") print(f" 等级: {result['risk_level']}") @@ -93,22 +117,24 @@ def main(): print(f" 回撤: {indicators.get('current_drawdown', 0):.1f}%") for a in result.get("alerts", []): print(f" ⚠️ {a}") + return 0 - elif cmd == "research": + elif args.cmd == "research": result = orch.run_research_cycle() top = result.get("research", {}).get("top_factors", []) print(f"\n因子评估结果:") if not top: print(" (无结果)") - return + return 0 print(f" {'因子':20s} {'IC':>8s} {'IC_IR':>8s} {'多头':>8s} {'空头':>8s} {'得分':>8s}") print(f" {'─'*60}") for f in top: print(f" {f['name']:20s} {f['ic_mean']:>+8.4f} {f['icir']:>8.3f} " f"{f['long_ret']:>+7.1f}% {f['short_ret']:>+7.1f}% {f['score']:>8.4f}") + return 0 - elif cmd == "warmup": - batch_n = int(sys.argv[2]) if len(sys.argv) > 2 else 50 + elif args.cmd == "warmup": + batch_n = args.batch_n print("首次批量预热: 每次 {} 只股票,分批执行...".format(batch_n)) sent = engines.get("sent") dm = engines.get("dm") @@ -129,24 +155,20 @@ def main(): print(" {} 失败: {}".format(ts_code, e)) print(" 累计同步: {} 条".format(total_synced)) print("预热完成: {} 条数据, {} 只新股票已缓存".format(total_synced, len(uncached))) + return 0 - elif cmd == "report": - date = sys.argv[2] if len(sys.argv) > 2 else None + elif args.cmd == "report": + date = args.date or today result = orch.generate_report(date=date) print("\n日报已生成: {}".format(result.get("report_path", "?"))) - # 存入 DB md = result.get("report_markdown", "") - if md: - from reports.storage import save_report - save_report(md, "量化日报", report_date=date or datetime.now().strftime("%Y%m%d"), - subject_type="daily", subject_code="") - print(" 已存入 DB") + # 注意: ReportAgent.execute 内部已保存到 DB,这里不再重复入库 if md: print(md) + return 0 - else: - print(f"未知命令: {cmd}") + return 0 # 防御:未匹配(应不会到达) if __name__ == "__main__": - main() + sys.exit(main()) \ No newline at end of file diff --git a/finance/cli/demo_ml.py b/finance/cli/demo_ml.py index ea6affb..e62aa5d 100644 --- a/finance/cli/demo_ml.py +++ b/finance/cli/demo_ml.py @@ -89,7 +89,7 @@ def main(): test_price = price_df.loc[X_test.index] test_factor = factor_df.loc[X_test.index] bt_engine = VectorBTEngine() - benchmark = MLBenchmark([lgb_model, cb_model], fe, test_price, test_factor, bt_engine) + benchmark = MLBenchmark([lgb_model, cb_model], fe, test_factor, test_price, bt_engine) result = benchmark.run() print(result.round(2).to_string()) print(" > 解读: 回测结果反映 ML 策略在测试集上的实盘表现。") diff --git a/finance/config/settings.py b/finance/config/settings.py index 45549a1..2c6adc5 100644 --- a/finance/config/settings.py +++ b/finance/config/settings.py @@ -44,3 +44,45 @@ AKSHARE_CONFIG = { # ── 默认参数 ────────────────────────────────────────────── DEFAULT_START_DATE = "20200101" DEFAULT_END_DATE = None # None 表示当天 + +# ── 数据新鲜度 ──────────────────────────────────────────── +# DB 窗口尾部落后于请求日多少个自然日,视为存在缺口需补拉(get_daily) +DATA_FRESHNESS_GAP_DAYS = int(os.getenv("DATA_FRESHNESS_GAP_DAYS", "7")) + +# ── 风险评估参数(RiskAgent) ───────────────────────────── +# 风险等级阈值:年化波动率(%) / 当前回撤(%),超过即升级 +RISK_THRESHOLDS = { + "high": {"vol": float(os.getenv("RISK_HIGH_VOL", "35")), "dd": float(os.getenv("RISK_HIGH_DD", "-15"))}, + "medium": {"vol": float(os.getenv("RISK_MED_VOL", "25")), "dd": float(os.getenv("RISK_MED_DD", "-8"))}, +} +# 各等级对应的建议总仓位、单票上限、止损线 +RISK_EXPOSURE = {"low": 0.85, "medium": 0.60, "high": 0.30} +RISK_MAX_SINGLE = {"low": 0.15, "medium": 0.10, "high": 0.05} +RISK_STOP_LOSS = {"low": -0.05, "medium": -0.08, "high": -0.12} + +# ── 选股参数(SelectionAgent) ──────────────────────────── +# 默认选股池大小(按股票列表前 N 只作为代表股池) +SELECTION_UNIVERSE_SIZE = int(os.getenv("SELECTION_UNIVERSE_SIZE", "100")) +# 单次打分上限(防止单次过慢) +SELECTION_SCORE_LIMIT = int(os.getenv("SELECTION_SCORE_LIMIT", "300")) +# 默认返回 Top N +SELECTION_TOP_N = int(os.getenv("SELECTION_TOP_N", "15")) +# 选股核心因子(覆盖多维度,减少计算量) +SELECTION_CORE_FACTORS = [ + "momentum_20", "momentum_60", + "rsi_14", "volatility_20", + "vol_ratio_5", "ma_dev_20", + "turnover_5", "amplitude_5", +] +# z-score 打分时是否做去极值(避免单股离群值主导 Top 排名) +SELECTION_WINSORIZE_ZSCORE = os.getenv("SELECTION_WINSORIZE_ZSCORE", "1") == "1" + +# ── 报告参数(ReportAgent) ──────────────────────────────── +# 日报市场概览指数代码 +REPORT_INDEX_CODES = { + "000001.SH": "上证指数", + "399001.SZ": "深证成指", + "399006.SZ": "创业板指", +} +# 日报标题中展示的 Top N(与 SELECTION_TOP_N 联动) +REPORT_TOP_N = SELECTION_TOP_N diff --git a/finance/data/data_manager.py b/finance/data/data_manager.py index 039d337..16a6e2f 100644 --- a/finance/data/data_manager.py +++ b/finance/data/data_manager.py @@ -5,6 +5,7 @@ DataManager — 统一数据管理层。 优先从 DB 读取,缺失时依次尝试 AkShare → Tushare 拉取并入库。 """ +import logging import time import pandas as pd @@ -15,6 +16,27 @@ from data.sources.tushare_source import TushareSource from database import dao from database.models import create_all_tables +logger = logging.getLogger(__name__) + + +def _rollback_days(yyyymmdd: str, days: int) -> str: + """把 'YYYYMMDD' 往前回退 days 个自然日,返回 'YYYYMMDD'。""" + from datetime import datetime, timedelta + dt = datetime.strptime(yyyymmdd, "%Y%m%d") - timedelta(days=days) + return dt.strftime("%Y%m%d") + + +def _safe_save_daily(df: pd.DataFrame, skip_exc: bool = True) -> None: + """把拉取到的日线入库;失败时记日志(可选抛错)。""" + cols = [c for c in dao._DAILY_COLS if c in df.columns] + try: + dao.save_daily(df[cols]) + except Exception as e: + if skip_exc: + logger.exception("[DataManager] 日线入库失败(数据已获取,未持久化):%s", e) + else: + raise + class DataManager: """统一数据管理。双数据源:AkShare(主)+ Tushare(备)。""" @@ -68,7 +90,7 @@ class DataManager: if df is not None and not df.empty: return df, label except Exception as e: - print(" [{}] {} 失败: {}".format(label, method_name, e)) + logger.warning(" [%s] %s 失败: %s", label, method_name, e) return pd.DataFrame(), None # ── 股票列表 ────────────────────────────────────────── @@ -77,15 +99,15 @@ class DataManager: if not force_refresh: df = dao.query_stock_list() if not df.empty: - print("[DataManager] 从 DB 读取股票列表: {} 只".format(len(df))) + logger.info("[DataManager] 从 DB 读取股票列表: %s 只", len(df)) return df - print("[DataManager] 拉取股票列表 (AkShare → Tushare)...") + logger.info("[DataManager] 拉取股票列表 (AkShare → Tushare)...") df, src = self._try_fetch("fetch_stock_list") if df.empty: - print("[DataManager] 所有数据源均无法获取股票列表") + logger.warning("[DataManager] 所有数据源均无法获取股票列表") return pd.DataFrame() - print("[DataManager] 股票列表已入库 ({}): {} 只".format(src, len(df))) + logger.info("[DataManager] 股票列表已入库 (%s): %s 只", src, len(df)) dao.save_stock_list(df) time.sleep(2) return df @@ -105,40 +127,48 @@ class DataManager: if not force_refresh: df = dao.query_daily(ts_code, start, end) if not df.empty: + # 检查 DB 窗口尾部是否明显落后于请求的 end(>5 个自然日)。 + # 若是,说明存在数据缺口,触发一次增量补拉,而不是把陈旧数据当完整返回。 + latest = str(df["trade_date"].astype(str).max()) + if end >= "20200101" and latest < _rollback_days(end, 7): + logger.warning( + "%s DB 窗口尾部 %s 明显落后于请求日 %s,触发补拉", ts_code, latest, end) + refreshed, src = self._try_fetch("fetch_daily", ts_code, latest, end) + if not refreshed.empty and src: + _safe_save_daily(refreshed) return df # DB 未命中 → 从数据源拉取 df, src = self._try_fetch("fetch_daily", ts_code, start, end) if df.empty: - print("[WARN] {} 日线获取失败 (AkShare+Tushare 均不可用)".format(ts_code)) + logger.warning("%s 日线获取失败 (AkShare+Tushare 均不可用)", ts_code) return pd.DataFrame() - # 保存到 DB - try: - # 筛选 DB 需要的列 - cols = [c for c in dao._DAILY_COLS if c in df.columns] - dao.save_daily(df[cols]) - except Exception as e: - print("[WARN] 日线入库失败: {}".format(e)) - + # 保存到 DB(失败仅记日志,不阻断已获取到的数据返回) + _safe_save_daily(df) return df def sync_daily(self, ts_code: str) -> int: latest = dao.get_latest_trade_date(ts_code) today = time.strftime("%Y%m%d") if latest and latest >= today: - print("[DataManager] {} 数据已是最新 ({})".format(ts_code, latest)) + logger.info("[DataManager] %s 数据已是最新 (%s)", ts_code, latest) return 0 start = latest or DEFAULT_START_DATE df, src = self._try_fetch("fetch_daily", ts_code, start, today) if df.empty: - print("[WARN] {} sync_daily 失败 (AkShare+Tushare 均不可用)".format(ts_code)) + logger.warning("[DataManager] %s sync_daily 失败 (AkShare+Tushare 均不可用)", ts_code) return 0 + # 入库失败要能反映到返回值,否则调用方会误以为已持久化 cols = [c for c in dao._DAILY_COLS if c in df.columns] - dao.save_daily(df[cols]) - print("[DataManager] {} 同步 {} 条日线 ({})".format(ts_code, len(df), src)) + try: + dao.save_daily(df[cols]) + except Exception as e: + logger.exception("[DataManager] %s 日线入库失败:%s", ts_code, e) + raise + logger.info("[DataManager] %s 同步 %s 条日线 (%s)", ts_code, len(df), src) return len(df) def sync_all_daily(self) -> int: @@ -148,11 +178,11 @@ class DataManager: try: total += self.sync_daily(ts_code) if (i + 1) % 50 == 0: - print("[DataManager] 进度: {}/{}".format(i + 1, len(stock_list))) + logger.info("[DataManager] 进度: %s/%s", i + 1, len(stock_list)) time.sleep(1) except Exception as e: - print("[WARN] {} 同步失败: {}".format(ts_code, e)) - print("[DataManager] 全量同步完成,新增 {} 条".format(total)) + logger.warning("[DataManager] %s 同步失败: %s", ts_code, e) + logger.info("[DataManager] 全量同步完成,新增 %s 条", total) return total # ── 财务数据 ────────────────────────────────────────── @@ -167,5 +197,5 @@ class DataManager: cols = [c for c in dao._FINA_COLS if c in df.columns] dao.save_financial(df[cols]) except Exception as e: - print("[WARN] 财务数据入库失败: {}".format(e)) + logger.exception("财务数据入库失败: %s", e) return df diff --git a/finance/database/dao.py b/finance/database/dao.py index b4e24d0..7a091a9 100644 --- a/finance/database/dao.py +++ b/finance/database/dao.py @@ -4,12 +4,16 @@ 提供 DataFrame 级别的读写操作,屏蔽底层 ORM/SQL 细节。 """ +import logging + import pandas as pd from sqlalchemy import text from database.connection import get_engine from database.models import StockBasic, StockDaily, StockFinancial, Report +logger = logging.getLogger(__name__) + # DB 表列名,供 DataManager 在写入前筛选 _DAILY_COLS = [ "ts_code", "trade_date", "open", "high", "low", "close", @@ -22,33 +26,59 @@ _FINA_COLS = [ ] -def _df_to_db(df: pd.DataFrame, model_class, replace: bool = False) -> int: - """将 DataFrame 写入对应表,返回写入行数。""" - if df.empty: +def _upsert_df(engine, model_class, df: pd.DataFrame) -> int: + """ + 用 SQLAlchemy Core 做 upsert(INSERT ... ON DUPLICATE KEY UPDATE)写库。 + + 不会 DROP/重建表(区别于 pandas to_sql 的 if_exists="replace"), + 主键冲突时按非主键列更新而非报错。返回受影响行数。 + + 说明: 需要真实 ORM model 的 __table__(含主键信息), + 因此不能像旧版那样用字符串表名 + pandas to_sql。 + """ + from sqlalchemy.dialects.mysql import insert + + if df.empty or len(df.columns) == 0: return 0 - engine = get_engine() - if_action = "replace" if replace else "append" - # 统一字符串列,避免 MySQL 类型问题 - df = df.where(pd.notna(df), None) - rows = len(df) - df.to_sql( - model_class.__tablename__, - con=engine, - if_exists=if_action, - index=False, - method="multi", - chunksize=500, - ) - return rows + table = model_class.__table__ + # 只保留表里真实存在的列,避免写入表外列导致 SQL 失败 + existing = [c.name for c in table.columns] + df = df[[c for c in df.columns if c in existing]].copy() + if df.empty or len(df.columns) == 0: + return 0 + # 统一 NaN -> None,交由 DB 处理;避免 pandas NA 类型报错 + data = df.where(pd.notna(df), None).to_dict(orient="records") + + pk_cols = [c.name for c in table.primary_key.columns] + non_pk = [c for c in df.columns if c in existing and c not in pk_cols] + if not non_pk: + # 只有主键列:用 INSERT IGNORE + stmt = insert(table).values(data).prefix_with("IGNORE") + else: + stmt = insert(table).values(data) + stmt = stmt.on_duplicate_key_update( + **{c: getattr(stmt.inserted, c) for c in non_pk} + ) + + with engine.begin() as conn: + result = conn.execute(stmt) + return result.rowcount # ── StockBasic ───────────────────────────────────────────── def save_stock_list(df: pd.DataFrame) -> int: - """保存股票列表(replace 模式)。""" + """ + 保存股票列表(upsert 模式,避免 replace 整表重建)。 + + 以 ts_code 为主键合并更新:新股票插入、已有股票按最新信息覆盖。 + """ cols = ["ts_code", "name", "area", "industry", "market", "list_date", "is_hs"] df = df[[c for c in cols if c in df.columns]].copy() - return _df_to_db(df, StockBasic, replace=True) + if df.empty: + return 0 + engine = get_engine() + return _upsert_df(engine, StockBasic, df) def query_stock_list() -> pd.DataFrame: @@ -60,7 +90,12 @@ def query_stock_list() -> pd.DataFrame: # ── StockDaily ───────────────────────────────────────────── def save_daily(df: pd.DataFrame) -> int: - """批量写入日线数据。先删旧再插新,避免主键冲突。""" + """ + 批量写入日线数据(单个事务内的原子 upsert)。 + + 复合主键 (ts_code, trade_date) 冲突时更新非主键列, + 因此修正后的历史 K 线会自动覆盖旧值,既不会整表重建也不会丢旧数据。 + """ cols = [ "ts_code", "trade_date", "open", "high", "low", "close", "pre_close", "change", "pct_chg", "vol", "amount", "turnover_rate", @@ -68,19 +103,8 @@ def save_daily(df: pd.DataFrame) -> int: df = df[[c for c in cols if c in df.columns]].copy() if df.empty: return 0 - # 删除即将写入的日期的旧数据 engine = get_engine() - ts_codes = df["ts_code"].unique().tolist() - trade_dates = df["trade_date"].unique().tolist() - if ts_codes and trade_dates: - with engine.connect() as conn: - conn.execute( - text("DELETE FROM {} WHERE ts_code IN :codes AND trade_date IN :dates".format( - StockDaily.__tablename__)), - {"codes": tuple(ts_codes), "dates": tuple(trade_dates)}, - ) - conn.commit() - return _df_to_db(df, StockDaily, replace=False) + return _upsert_df(engine, StockDaily, df) def query_daily(ts_code: str, start: str | None = None, end: str | None = None) -> pd.DataFrame: @@ -115,14 +139,22 @@ def get_latest_trade_date(ts_code: str) -> str | None: # ── StockFinancial ───────────────────────────────────────── def save_financial(df: pd.DataFrame) -> int: - """批量写入财务数据(replace 模式:同报告期覆盖更新)。""" + """ + 批量写入财务数据(upsert 模式,同报告期 (ts_code, end_date) 覆盖更新)。 + + 关键: 不使用 pandas to_sql 的 if_exists="replace"(那会 DROP 整表重建, + 导致逐股写入时清空所有其他股票的财务记录并丢失 ORM 主键/索引)。 + """ cols = [ "ts_code", "end_date", "eps", "bvps", "roe", "roe_diluted", "net_profit_margin", "debt_to_assets", "current_ratio", "quick_ratio", "total_revenue", "total_revenue_yoy", "net_profit", "net_profit_yoy", ] df = df[[c for c in cols if c in df.columns]].copy() - return _df_to_db(df, StockFinancial, replace=True) + if df.empty: + return 0 + engine = get_engine() + return _upsert_df(engine, StockFinancial, df) def query_financial(ts_code: str) -> pd.DataFrame: diff --git a/finance/factors/fundamental/_mapping.py b/finance/factors/fundamental/_mapping.py new file mode 100644 index 0000000..2471512 --- /dev/null +++ b/finance/factors/fundamental/_mapping.py @@ -0,0 +1,100 @@ +""" +财务数据 → 日线映射(负责消除披露时点前视偏差)。 + +A 股季报的披露日远晚于报告期末: + - 一季报 / 年报 最迟约 4/30 + - 中报 最迟约 8/31 + - 三季报 最迟约 10/31 + +若直接把报告期 end_date 当日就"看到"本期财务结果,会引入前视偏差。 +本模块统一在 end_date 上叠加一个保守的披露滞后: + 1) 财务帧中若带 ann_date(实际披露日),则优先用 ann_date 作为可用日; + 2) 否则按报告期月份推断法定披露时点,作为保守可用日。 +""" + +from __future__ import annotations + +import pandas as pd + + +def _disclosure_available_date(end_date_ts: pd.Timestamp) -> pd.Timestamp: + """根据报告期末推断法定披露可用日(无 ann_date 时的保守近似)。 + + 一季度(0331)→4/30;中报(0630)→8/31;三季报(0930)→10/31;年报(1231)→次年4/30。 + """ + if end_date_ts.month == 3 and end_date_ts.day == 31: + return pd.Timestamp(year=end_date_ts.year, month=4, day=30) + if end_date_ts.month == 6 and end_date_ts.day == 30: + return pd.Timestamp(year=end_date_ts.year, month=8, day=31) + if end_date_ts.month == 9 and end_date_ts.day == 30: + return pd.Timestamp(year=end_date_ts.year, month=10, day=31) + # 年报 12/31 → 次年 4/30 + return pd.Timestamp(year=end_date_ts.year + 1, month=4, day=30) + + +def _eff_available_dates(fina_df: pd.DataFrame) -> pd.DataFrame: + """计算每期财务的"可用日"(取最小滞后:有 ann_date 用它,否则法定截止)。""" + fina = fina_df.copy() + # 数值型 ann_date(YYYYMMDD)→ datetime;缺失用报告期末推断 + if "ann_date" in fina.columns: + ann = pd.to_datetime(fina["ann_date"].astype(str), format="%Y%m%d", errors="coerce") + else: + ann = pd.Series(pd.NaT, index=fina.index) + + end = pd.to_datetime( + fina["end_date"].astype(str), format="%Y%m%d", errors="coerce") + + avail = ann.fillna(pd.Series( + [_disclosure_available_date(x) if not pd.isna(x) else pd.NaT for x in end], + index=end.index, + )) + # 极少数 ann_date 早于报告期末(脏数据)时兜底用期末 + avail = avail.where(avail >= end, end) + fina["_avail"] = avail + return fina + + +def effective_available_dates(fina_df: pd.DataFrame) -> pd.DataFrame: + """ + (公开) 返回带 _avail(披露可用日)的财务帧,供同比/环比等因子复用。 + 要求 fina_df 至少含 'end_date';可选 'ann_date'。 + """ + return _eff_available_dates(fina_df) + + +def map_fundamental_to_daily( + daily_df: pd.DataFrame, + fina_df: pd.DataFrame, + column: str, +) -> pd.Series: + """ + 将季度财务数据按披露可用日映射到日线索引(前值填充)。 + + - 优先按 ann_date(实际披露日)对齐; + - 无 ann_date 时按报告期末的法定披露截止日保守对齐; + - 从而避免"财报在披露日之前就被回测看到"的前视偏差。 + """ + if column not in fina_df.columns or "end_date" not in fina_df.columns: + return pd.Series(float("nan"), index=daily_df.index) + + cols = ["end_date", column] + (["ann_date"] if "ann_date" in fina_df.columns else []) + fina = fina_df[cols].dropna(subset=["end_date", column]).copy() + if fina.empty: + return pd.Series(float("nan"), index=daily_df.index) + + fina = _eff_available_dates(fina) + fina = fina.sort_values("_avail") + + dates = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce") + result = pd.Series(float("nan"), index=daily_df.index) + avail_dates = fina["_avail"].values + + for i, avail_dt in enumerate(avail_dates): + if pd.isna(avail_dt): + continue + mask = dates >= avail_dt + if i + 1 < len(avail_dates) and not pd.isna(avail_dates[i + 1]): + mask &= dates < avail_dates[i + 1] + result[mask] = fina[column].iloc[i] + + return result.astype("float64") \ No newline at end of file diff --git a/finance/factors/fundamental/pe_pb.py b/finance/factors/fundamental/pe_pb.py index e972d4a..2587984 100644 --- a/finance/factors/fundamental/pe_pb.py +++ b/finance/factors/fundamental/pe_pb.py @@ -7,6 +7,7 @@ PE / PB 估值因子。 import pandas as pd from factors.base import BaseFactor +from factors.fundamental._mapping import map_fundamental_to_daily class PEFactor(BaseFactor): @@ -31,7 +32,7 @@ class PEFactor(BaseFactor): if self._financial_df is None or self._financial_df.empty: return pd.Series(float("nan"), index=df.index) - eps_series = _map_to_daily(df, self._financial_df, "eps") + eps_series = map_fundamental_to_daily(df, self._financial_df, "eps") close = df["close"] return close / eps_series.replace(0, float("nan")) @@ -56,7 +57,7 @@ class PBFactor(BaseFactor): if self._financial_df is None or self._financial_df.empty: return pd.Series(float("nan"), index=df.index) - bvps_series = _map_to_daily(df, self._financial_df, "bvps") + bvps_series = map_fundamental_to_daily(df, self._financial_df, "bvps") return df["close"] / bvps_series.replace(0, float("nan")) def get_required_columns(self) -> list[str]: @@ -80,34 +81,8 @@ class EPFactor(BaseFactor): if self._financial_df is None or self._financial_df.empty: return pd.Series(float("nan"), index=df.index) - eps_series = _map_to_daily(df, self._financial_df, "eps") + eps_series = map_fundamental_to_daily(df, self._financial_df, "eps") return eps_series / df["close"].replace(0, float("nan")) * 100 def get_required_columns(self) -> list[str]: return ["close"] - - -def _map_to_daily( - daily_df: pd.DataFrame, - fina_df: pd.DataFrame, - column: str, -) -> pd.Series: - """将季度财务数据填充到日线索引(前值填充)。""" - fina = fina_df[["end_date", column]].dropna().copy() - fina["end_date"] = fina["end_date"].astype(str) - fina = fina.sort_values("end_date") - - result = pd.Series(float("nan"), index=daily_df.index) - if fina.empty: - return result - - dates = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce") - fina_dates = pd.to_datetime(fina["end_date"], format="%Y%m%d", errors="coerce") - - for i, fina_date in enumerate(fina_dates): - mask = dates >= fina_date - if i + 1 < len(fina_dates): - mask &= dates < fina_dates.iloc[i + 1] - result[mask] = fina[column].iloc[i] - - return result.astype("float64") diff --git a/finance/factors/fundamental/roe.py b/finance/factors/fundamental/roe.py index 391903b..950c4fe 100644 --- a/finance/factors/fundamental/roe.py +++ b/finance/factors/fundamental/roe.py @@ -5,6 +5,7 @@ ROE 因子。 import pandas as pd from factors.base import BaseFactor +from factors.fundamental._mapping import map_fundamental_to_daily, effective_available_dates class ROEFactor(BaseFactor): @@ -45,28 +46,8 @@ class ROEFactor(BaseFactor): fina_df: pd.DataFrame, column: str, ) -> pd.Series: - """将财务数据(季度)映射到日线索引。""" - fina = fina_df[["end_date", column]].dropna().copy() - fina["end_date"] = fina["end_date"].astype(str) - fina = fina.sort_values("end_date") - - result = pd.Series(float("nan"), index=daily_df.index) - - if fina.empty: - return result - - dates = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce") - fina_dates = pd.to_datetime(fina["end_date"], format="%Y%m%d", errors="coerce") - - for i, fina_date in enumerate(fina_dates): - mask = dates >= fina_date - if i + 1 < len(fina_dates): - mask &= dates < fina_dates.iloc[i + 1] - else: - pass # 最新一期覆盖所有后续日期 - result[mask] = fina[column].iloc[i] - - return result.astype("float64") + """将财务数据(季度)映射到日线索引,按披露可用日消除前视。""" + return map_fundamental_to_daily(daily_df, fina_df, column) class ROETTMDeltaFactor(BaseFactor): @@ -82,23 +63,36 @@ class ROETTMDeltaFactor(BaseFactor): if self._financial_df is None or self._financial_df.empty: return pd.Series(float("nan"), index=df.index) - fina = self._financial_df[["end_date", "roe"]].dropna().copy() - fina["end_date"] = fina["end_date"].astype(str) - fina["year"] = fina["end_date"].str[:4].astype(int) - fina = fina.sort_values("end_date") + if "roe" not in self._financial_df.columns or "end_date" not in self._financial_df.columns: + return pd.Series(float("nan"), index=df.index) + + # 基于财务帧计算披露可用日(含 ann_date 优先/法定滞后兜底) + fina = effective_available_dates(self._financial_df) + fina = fina[["_avail", "end_date", "roe"]].dropna(subset=["_avail", "roe"]).copy() + if fina.empty: + return pd.Series(float("nan"), index=df.index) + + # 报告期月份(A 股季末:03/06/09/12) + end_dt = pd.to_datetime(fina["end_date"].astype(str), format="%Y%m%d", errors="coerce") + fina["report_month"] = end_dt.dt.month + fina["report_year"] = end_dt.dt.year + # 本年取数映射:{ (year, month): roe } + cur_map = dict(zip(zip(fina["report_year"], fina["report_month"]), fina["roe"])) + # 按披露可用日排序,逐期覆盖区间 + fina = fina.sort_values("_avail") - # 按年分组计算 YoY 差值 roe_delta = pd.Series(float("nan"), index=df.index) dates = pd.to_datetime(df.index, format="%Y%m%d", errors="coerce") for _, row in fina.iterrows(): - this_year = row["year"] - prev_row = fina[fina["year"] == this_year - 1] - if prev_row.empty: + avail_dt = row["_avail"] + if pd.isna(avail_dt): continue - delta = row["roe"] - prev_row["roe"].iloc[-1] - f_date = pd.to_datetime(row["end_date"], format="%Y%m%d") - mask = dates >= f_date + prev_roe = cur_map.get((row["report_year"] - 1, row["report_month"])) + if prev_roe is None: + continue + delta = row["roe"] - prev_roe + mask = dates >= avail_dt roe_delta[mask] = delta return roe_delta.astype("float64") diff --git a/finance/factors/registry.py b/finance/factors/registry.py index ab37f88..439c55f 100644 --- a/finance/factors/registry.py +++ b/finance/factors/registry.py @@ -101,6 +101,10 @@ def get_factor(name: str, **overrides) -> BaseFactor: 返回: BaseFactor 实例 + + 说明: 返回实例的 .name 恒等于注册键 name,即使构造器默认生成的 + .name 与注册键不同(如 boll 的构造器默认 .name='boll_20')—— + 注册键是列名的唯一事实源,避免 FactorEngine.compute 的列名漂移。 """ if name not in _BUILTIN_FACTORIES: raise KeyError(f"未知因子: '{name}'。可用: {list(_BUILTIN_FACTORIES)}") @@ -109,9 +113,9 @@ def get_factor(name: str, **overrides) -> BaseFactor: for k, v in overrides.items(): if hasattr(factor, k): setattr(factor, k, v) - # 更新 factor.name - if hasattr(factor, "name"): - factor.name = name + # 无论是否覆盖参数,都强制 .name = 注册键,保证与分类/策略的引用一致 + if hasattr(factor, "name"): + factor.name = name return factor diff --git a/finance/factors/sentiment/news_source.py b/finance/factors/sentiment/news_source.py index cc75852..a387428 100644 --- a/finance/factors/sentiment/news_source.py +++ b/finance/factors/sentiment/news_source.py @@ -10,6 +10,7 @@ """ import json +import logging import os import time from datetime import datetime, timedelta @@ -21,6 +22,8 @@ import requests # 确保 .env 已加载 import config.settings # noqa: F401 +logger = logging.getLogger(__name__) + class NewsSource: """ @@ -87,7 +90,7 @@ class NewsSource: frames.append(df) time.sleep(self.request_delay) except Exception as e: - print(" [WARN] AkShare 新闻获取失败 ({}): {}".format(ts_code, e)) + logger.warning("[news_source] AkShare 新闻获取失败 (%s): %s", ts_code, e) if self.use_xwlb: try: @@ -99,7 +102,7 @@ class NewsSource: if not df.empty: frames.append(df) except Exception as e: - print(" [WARN] xwlb 新闻获取失败: {}".format(e)) + logger.warning("[news_source] xwlb 新闻获取失败: %s", e) if self.use_mcp: try: @@ -107,7 +110,7 @@ class NewsSource: if not df.empty: frames.append(df) except Exception as e: - print(" [WARN] MCP 新闻获取失败: {}".format(e)) + logger.warning("[news_source] MCP 新闻获取失败: %s", e) if not frames: return pd.DataFrame(columns=["date", "title", "content", "source", "url"]) @@ -133,7 +136,8 @@ class NewsSource: symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "") try: df = ak.stock_news_em(symbol=symbol.zfill(6)) - except Exception: + except Exception as e: + logger.warning("[news_source] AkShare 新闻接口失败 (%s): %s", ts_code, e) return pd.DataFrame() if df is None or df.empty: @@ -204,7 +208,8 @@ class NewsSource: df["url"] = "" return df[["date", "title", "content", "source", "url"]] - except Exception: + except Exception as e: + logger.warning("[news_source] xwlb 新闻获取失败: %s", e) return pd.DataFrame() # ── MCP 数据源 ──────────────────────────────────────── @@ -232,7 +237,8 @@ class NewsSource: df = pd.DataFrame(records) df["source"] = "mcp_trendradar" return df[["date", "title", "content", "source", "url"]] - except Exception: + except Exception as e: + logger.warning("[news_source] MCP 新闻获取失败: %s", e) return pd.DataFrame() def _mcp_initialize(self) -> str | None: @@ -267,10 +273,10 @@ class NewsSource: if session_id: self._mcp_session_id = session_id else: - print(" [WARN] MCP initialize 未返回 session-id") + logger.warning("[news_source] MCP initialize 未返回 session-id") return session_id except Exception as e: - print(" [WARN] MCP 连接失败: {}".format(e)) + logger.warning("[news_source] MCP 连接失败: %s", e) return None def _mcp_call_tool( @@ -299,7 +305,8 @@ class NewsSource: data = json.loads(line[5:].strip()) return data.get("result", {}) return None - except Exception: + except Exception as e: + logger.warning("[news_source] MCP call_tool 失败: %s", e) return None @staticmethod diff --git a/finance/factors/sentiment/sentiment_engine.py b/finance/factors/sentiment/sentiment_engine.py index 18fdb3e..4d53bdb 100644 --- a/finance/factors/sentiment/sentiment_engine.py +++ b/finance/factors/sentiment/sentiment_engine.py @@ -128,11 +128,12 @@ class SentimentEngine: # 5. Qwen 情绪分析(有 API key 时才执行) sentiment_df = self._analyze_news(news_df) - # 6. 计算因子 + # 6. 计算因子(对齐注册表,产出全部 4 个注册情绪因子,含 news_sent_20) factor_dfs = {} if not sentiment_df.empty: for factor_cls, kwargs in [ (NewsSentimentFactor, {"window": 5, "sentiment_df": sentiment_df}), + (NewsSentimentFactor, {"window": 20, "sentiment_df": sentiment_df}), (SentimentConfidenceFactor, {"window": 5, "sentiment_df": sentiment_df}), (SentimentMomentumFactor, {"period": 5, "sentiment_df": sentiment_df}), ]: diff --git a/finance/factors/technical/rsi.py b/finance/factors/technical/rsi.py index 12bff9f..bbc0612 100644 --- a/finance/factors/technical/rsi.py +++ b/finance/factors/technical/rsi.py @@ -2,6 +2,7 @@ RSI 相对强弱因子。 """ +import numpy as np import pandas as pd from factors.base import BaseFactor @@ -22,8 +23,13 @@ class RSIFactor(BaseFactor): loss = (-delta).clip(lower=0) avg_gain = gain.ewm(span=self.period, min_periods=self.period).mean() avg_loss = loss.ewm(span=self.period, min_periods=self.period).mean() - rs = avg_gain / avg_loss.replace(0, float("nan")) - return 100 - 100 / (1 + rs) + # 标准 Wilder RSI:avg_loss==0 时 RSI 应 = 100,而非 NaN。 + # 用 where 显式处理除零,避免 replace(0, nan) 把上涨趋势判为缺失。 + rs = avg_gain / avg_loss.where(avg_loss != 0, np.nan) + rsi = 100 - 100 / (1 + rs) + # 上涨且无下跌的高位情形补 100(无 prior-loss 的窗口仍留 NaN 由上游填充) + rsi = rsi.where(avg_loss != 0, 100.0) + return rsi def get_required_columns(self) -> list[str]: return ["close"] diff --git a/finance/models/backtest_integration.py b/finance/models/backtest_integration.py index fee2439..38d72c3 100644 --- a/finance/models/backtest_integration.py +++ b/finance/models/backtest_integration.py @@ -82,34 +82,46 @@ class MLStrategy(BaseStrategy): class MLBenchmark: - """ML 模型基准对比测试。""" + """ML 模型基准对比测试。 + + 入参的 feature_engine 应已用模型训练集 fit 过(scaler/winsor/median 已缓存)。 + run() 对 **样本外测试数据**(test_factor_df/test_price_df)用 fit=False 转换后 + 计算预测 IC,避免"同一时间段既训练又评估"的前视泄漏。 + """ def __init__( self, models: list[BaseModel], feature_engine: FeatureEngine, - price_df: pd.DataFrame, - factor_df: pd.DataFrame, + test_factor_df: pd.DataFrame, + test_price_df: pd.DataFrame, bt_engine: VectorBTEngine | None = None, ): self.models = models self.feature_engine = feature_engine - self.price_df = price_df - self.factor_df = factor_df + self.test_factor_df = test_factor_df + self.test_price_df = test_price_df self.bt_engine = bt_engine or VectorBTEngine() def run(self) -> pd.DataFrame: - """对比各模型的预测质量和回测表现。""" + """对比各模型在样本外测试集上的预测质量和回测表现。""" rows = [] + # 标签由规则(前视收益)决定,预测时可直接用同一 build_labels 构造, + # 避免依赖 build(fit=False) 不产标签的语义。 + y_test = self.feature_engine.build_labels(self.test_price_df) + X_test, _ = self.feature_engine.build(self.test_factor_df, self.test_price_df, fit=False) + if X_test is None or X_test.empty or y_test is None or y_test.dropna().empty: + raise RuntimeError( + "MLBenchmark: 样本外测试集为空或 feature_engine 未在训练集上 fit。") + for model in self.models: + # OOS 回测:MLStrategy 用同一已 fit engine 对测试集生成信号 strategy = MLStrategy(model, self.feature_engine) - report = self.bt_engine.run(strategy, self.price_df, self.factor_df) + report = self.bt_engine.run(strategy, self.test_price_df, self.test_factor_df) - # OOS 预测 vs 真实值 - X, y_true = self.feature_engine.build(self.factor_df, self.price_df, fit=True) - y_pred = model.predict(X) - - ic = y_pred.corr(y_true) if len(y_pred) > 0 else 0 + preds = model.predict(X_test) + common = X_test.index.intersection(y_test.dropna().index) + ic = preds.reindex(common).astype(float).corr(y_test.reindex(common).astype(float)) if len(common) > 1 else 0 rows.append({ "model": model.name, diff --git a/finance/models/features.py b/finance/models/features.py index d8af9dd..aefc7fd 100644 --- a/finance/models/features.py +++ b/finance/models/features.py @@ -1,7 +1,9 @@ """ 特征工程:因子 → 特征矩阵 + 目标标签。 -严禁使用未来数据。所有变换基于 expanding window 或训练集统计。 +严禁使用未来数据。所有变换的统计量(去极值边界、NaN 填充中位数、缩放器) +只在 fit(训练)阶段从训练样本估算并缓存在 self 上,predict 阶段复用这些 +训练统计,避免训练/推理分布不一致(泄漏)和跨股票重复 refit scaler。 """ import numpy as np @@ -34,6 +36,20 @@ class FeatureEngine: self._scaler = RobustScaler() self._scaler_fitted = False self._valid_features: list[str] = [] + # 训练阶段缓存的统计量,predict 阶段复用 + self._winsor_lower: pd.Series = pd.Series(dtype=float) + self._winsor_upper: pd.Series = pd.Series(dtype=float) + self._fill_medians: pd.Series = pd.Series(dtype=float) + + def reset(self): + """清空训练状态,便于重新 fit 新的训练集。""" + self._scaler = RobustScaler() + self._scaler_fitted = False + self._valid_features = [] + self._winsor_lower = pd.Series(dtype=float) + self._winsor_upper = pd.Series(dtype=float) + self._fill_medians = pd.Series(dtype=float) + return self # ── 标签构建 ────────────────────────────────────────── @@ -49,10 +65,55 @@ class FeatureEngine: ret = (future - close) / close * 100 if self.label_type == "classification": - return (ret > 0).astype(int) + # 末尾 lookahead 行无法构建标签,用 NaN 标记而非强制判负(避免标签偏差) + cls = (ret > 0).astype(float) + cls = cls.where(~ret.isna(), np.nan) + return cls.rename(f"y_fwd_{self.lookahead}") return ret.rename(f"y_fwd_{self.lookahead}") + # ── 特征变换(fit 估算统计 / predict 复用统计) ──────── + + def _winsorize_bounds(self, X: pd.DataFrame): + lo, hi = self.winsorize_pct + q = X.quantile([lo, hi]) + self._winsor_lower = q.loc[lo] + self._winsor_upper = q.loc[hi] + + def _apply_transform(self, X: pd.DataFrame, fit: bool) -> pd.DataFrame: + """去极值 + NaN 填充 + 缩放。fit 时估算并缓存统计,否则复用。""" + X = X.copy() + + # NaN 填充:前值填充,缺失再按记录的中位数填充 + X = X.ffill() + if fit: + # 用有效特征(非全 NaN 列)做列中位数 + self._fill_medians = X.median() + for col in X.columns: + if col in self._fill_medians: + X[col] = X[col].fillna(self._fill_medians[col]) + + # 去极值 + if fit: + self._winsorize_bounds(X) + for col in X.columns: + if col in self._winsor_lower.index and col in self._winsor_upper.index: + X[col] = X[col].clip(self._winsor_lower[col], self._winsor_upper[col]) + + # 缩放:fit 时 fit_transform,predict 时 transform(复用训练统计) + if fit: + X_scaled = self._scaler.fit_transform(X) + self._scaler_fitted = True + else: + if not self._scaler_fitted: + raise RuntimeError( + "FeatureEngine 尚未 fit,无法在 predict 模式下 transform。" + "必须先用 fit=True 调用 build 训练缩放统计。" + ) + X_scaled = self._scaler.transform(X) + + return pd.DataFrame(X_scaled, index=X.index, columns=X.columns) + # ── 特征构建 ────────────────────────────────────────── def build( @@ -74,40 +135,28 @@ class FeatureEngine: """ X = factor_df.copy() - # 1. 剔除 NaN 率过高的列 + # 1. 剔除 NaN 率过高的列(只在 fit 时决定,predict 沿用同一列集) if fit: nan_ratio = X.isna().mean() - self._valid_features = list(nan_ratio[nan_ratio <= self.nan_threshold].index) - # 排除非因子列 - self._valid_features = [c for c in self._valid_features - if c not in ("close", "open", "high", "low", "volume")] - X = X[self._valid_features].copy() if self._valid_features else X + excluded = ("close", "open", "high", "low", "volume") + self._valid_features = [ + c for c in X.columns + if nan_ratio[c] <= self.nan_threshold and c not in excluded + ] + if not self._valid_features: + # 没有有效特征 → 空矩阵 + return pd.DataFrame(index=X.index), None + if not self._valid_features: + # predict 且从未 fit → 无有效特征 + return pd.DataFrame(index=X.index), None + X = X[self._valid_features].copy() - # 2. 缺失值填充:前值填充 → 截面中位数 - X = X.ffill().fillna(X.median()) + X = self._apply_transform(X, fit=fit) - # 3. 去极值(Winsorize) - if fit: - lo, hi = self.winsorize_pct - self._winsor_lower = X.quantile(lo) - self._winsor_upper = X.quantile(hi) - for col in X.columns: - if col in getattr(self, "_winsor_lower", pd.Series()): - X[col] = X[col].clip(self._winsor_lower[col], self._winsor_upper[col]) - - # 4. 标准化(训练时 fit,预测时 transform) - if fit: - X_scaled = self._scaler.fit_transform(X) - self._scaler_fitted = True - else: - X_scaled = self._scaler.transform(X) - - X = pd.DataFrame(X_scaled, index=X.index, columns=X.columns) - - # 5. 构建标签 + # 2. 构建标签 y = self.build_labels(price_df) if fit else None - # 6. 对齐(删掉无法构建标签的行) + # 3. 对齐(删掉无法构建标签的行) if fit: valid_idx = X.index.intersection(y.dropna().index) X = X.loc[valid_idx] @@ -115,28 +164,76 @@ class FeatureEngine: return X, y - # ── 多股票构建 ──────────────────────────────────────── + # ── 多股票构建(一次性 fit,消除跨股票 refit 泄漏) ──── def build_universe( self, factor_universe: dict[str, pd.DataFrame], price_universe: dict[str, pd.DataFrame], ) -> tuple[pd.DataFrame, pd.Series]: - """多股票拼接特征矩阵(每只股票独立处理再拼接)。""" - X_parts, y_parts = [], [] + """ + 多股票拼接特征矩阵。 + + 相比旧版(每只股票独立 fit=True 反复 refit scaler),现在: + - 先拼所有股票的因子值为一张横截面表,统一一次性 fit 缩放统计, + 保证跨股票同分布; + - 标签按每只股票自身的前向收益构建,避免未来的跨股票串档。 + """ + # 1. 收集每只股票有效期间内的特征行(保留 _ts_code 以区分) + parts: list[pd.DataFrame] = [] + key_order: list[str] = [] for ts_code in factor_universe: f_df = factor_universe[ts_code] p_df = price_universe.get(ts_code) - if p_df is None or f_df.empty or p_df.empty: + if p_df is None or f_df.empty or p_df.empty or "close" not in p_df.columns: continue - X, y = self.build(f_df, p_df, fit=True) - if X.empty: + common = f_df.index.intersection(p_df.index) + if len(common) == 0: continue - X["_ts_code"] = ts_code - X_parts.append(X) - y_parts.append(y) - if not X_parts: + f_df = f_df.loc[common] + f_df = f_df.copy() + f_df["_ts_code"] = ts_code + parts.append(f_df) + key_order.append(ts_code) + + if not parts: return pd.DataFrame(), pd.Series() - X_all = pd.concat(X_parts) - y_all = pd.concat(y_parts) - return X_all.drop(columns=["_ts_code"]), y_all + + X_all = pd.concat(parts) + + # 2. 剔除 NaN 率过高的列(基于全横截面 fit) + nan_ratio = X_all.isna().mean() + excluded = ("close", "open", "high", "low", "volume", "_ts_code") + self._valid_features = [ + c for c in X_all.columns + if nan_ratio[c] <= self.nan_threshold and c not in excluded + ] + if not self._valid_features: + return pd.DataFrame(), pd.Series() + + # 3. 一次性 fit 变换统计并应用(单次跨股票) + feat = X_all[self._valid_features].copy() + feat_scaled = self._apply_transform(feat, fit=True) + + # 4. 每只股票构建自身标签并对齐(不把标签跨股票串起来) + y_parts = [] + rows = [] + for ts_code in key_order: + rows_mask = X_all["_ts_code"] == ts_code + f_local = feat_scaled[rows_mask] + p_local = price_universe[ts_code].loc[f_local.index] + y_local = self.build_labels(p_local).dropna() + keep = f_local.index.intersection(y_local.index) + if len(keep) == 0: + continue + rows.append(f_local.loc[keep]) + y_parts.append(y_local.loc[keep]) + + if not rows: + return pd.DataFrame(), pd.Series() + + X_out = pd.concat(rows) + y_out = pd.concat(y_parts) + if "_ts_code" in X_out.columns: + X_out = X_out.drop(columns=["_ts_code"]) + return X_out, y_out \ No newline at end of file diff --git a/finance/optimizer/engine.py b/finance/optimizer/engine.py index 201f95b..d237c42 100644 --- a/finance/optimizer/engine.py +++ b/finance/optimizer/engine.py @@ -129,9 +129,13 @@ class OptunaEngine: n_total = len(price_df) windows = [] - test_equities = [] param_history = [] + # 跨窗口结转资本:consolidated 是连续可投资的净值曲线, + # 每个测试窗口的收益按期初已积累资本放大,而非各自从 100k 独立重算。 + running_capital = float(self.bt_engine.initial_capital) + equity_parts: list[pd.Series] = [] + start = 0 while start + train_window + test_window <= n_total: train_slice = slice(start, start + train_window) @@ -162,8 +166,14 @@ class OptunaEngine: test_report = self.bt_engine.run(test_strategy, test_price, test_factor) - if len(test_report.equity_curve) > 0: - test_equities.append(test_report.equity_curve) + # 该窗口的相对收益 → 用累计资本放大 → 连续资本曲线 + eq = test_report.equity_curve + if eq is not None and len(eq) > 0: + window_ret = eq.pct_change().fillna(0.0) + # 用上一窗口末累计资本作基准放大本窗口收益 + window_capital = running_capital * (1 + window_ret).cumprod() + equity_parts.append(window_capital) + running_capital = float(window_capital.iloc[-1]) train_idx = train_price.index test_idx = test_price.index @@ -182,8 +192,8 @@ class OptunaEngine: start += test_window - # 合并测试期权益曲线 - consolidated = _merge_test_periods(test_equities, self.bt_engine.initial_capital) + # 合并测试期权益曲线(跨窗口结转后的连续净值) + consolidated = _merge_test_periods(equity_parts) # 参数稳定性 param_df = pd.DataFrame(param_history) if param_history else pd.DataFrame() @@ -198,20 +208,20 @@ class OptunaEngine: def _merge_test_periods( - equity_list: list[pd.Series], - initial_capital: float = 100_000, + equity_parts: list[pd.Series], ) -> BacktestReport | None: - """拼接各窗口测试期权益曲线为一个连续序列。""" - if not equity_list: + """拼接已跨窗口结转资本的测试期净值片段为连续序列。""" + if not equity_parts: return None - merged = pd.concat(equity_list) - merged = merged.sort_index() - merged = merged[~merged.index.duplicated()] + merged = pd.concat(equity_parts) + merged = merged[~merged.index.duplicated(keep="last")].sort_index() # 确保 DatetimeIndex if not isinstance(merged.index, pd.DatetimeIndex): - merged.index = pd.to_datetime(merged.index, format="%Y%m%d") + parsed = pd.to_datetime(merged.index, format="%Y%m%d", errors="coerce") + if parsed.notna().all(): + merged.index = parsed dd = merged / merged.cummax() - 1 daily_ret = merged.pct_change().dropna() diff --git a/finance/tests/__init__.py b/finance/tests/__init__.py new file mode 100644 index 0000000..80f9295 --- /dev/null +++ b/finance/tests/__init__.py @@ -0,0 +1,6 @@ +""" +finance 子项目回归测试包。 + +运行方式(在 finance/ 下): + python -m pytest tests/ -v +""" \ No newline at end of file diff --git a/finance/tests/test_agents.py b/finance/tests/test_agents.py new file mode 100644 index 0000000..bff71dc --- /dev/null +++ b/finance/tests/test_agents.py @@ -0,0 +1,87 @@ +""" +Agent 层回归测试:RiskAgent 配置集中可用性 + Orchestrator 每日流程步骤隔离。 +""" +import sys +import os +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import numpy as np +import pandas as pd +import pytest +from unittest.mock import patch + +from agents.risk_agent import RiskAgent +from agents.orchestrator import AgentOrchestrator, TOTAL_STEPS + + +def _flat_dm(): + """给出平淡行情 → 风险 low。""" + class _DM: + def get_daily(self, code): + idx = pd.date_range('2025-01-01', periods=300, freq='B').strftime('%Y%m%d') + close = np.linspace(100, 100 + 0.03 * 300, 300) + np.random.default_rng(0).normal(0, 0.2, 300) + return pd.DataFrame({'trade_date': idx, 'close': close}) + return _DM() + + +def _hot_dm(): + class _DM: + def get_daily(self, code): + idx = pd.date_range('2025-01-01', periods=300, freq='B').strftime('%Y%m%d') + close = 100 * np.exp(np.cumsum(np.random.default_rng(1).normal(0, 0.03, 300))) + return pd.DataFrame({'trade_date': idx, 'close': close}) + return _DM() + + +def test_risk_config_low_level(): + r = RiskAgent(dm=_flat_dm()).execute() + assert r['risk_level'] in ('low', 'medium', 'high') + assert r['target_exposure'] in (0.85, 0.60, 0.30) + + +def test_risk_config_high_level(): + r = RiskAgent(dm=_hot_dm()).execute() + assert r['risk_level'] in ('high', 'medium') + + +class _FailAgent: + def execute(self, **kw): + raise RuntimeError("boom") + + +class _OkAgent: + def execute(self, **kw): + return {"report_path": "/tmp/x.md", "risk_level": "low", + "target_exposure": 0.8, "top_picks": [], "date": "20260601"} + + +class _FakeDM: + def get_stock_list(self): + return pd.DataFrame(index=['000001.SZ', '600519.SH']) + def sync_daily(self, code): + return 5 + + +class _FakeSent: + def get_scope_stocks(self): + return ["000001.SZ", "600519.SH"] + def compute(self, *a, **k): + return pd.DataFrame({"news_sent_5": [0.1, 0.2]}) + + +def test_run_daily_steps_isolated(): + """risk/selection 抛错被隔离,report/sentiment/sync 仍成功。""" + with patch('database.dao.get_latest_trade_date', side_effect=lambda c: "20260530"): + orch = AgentOrchestrator(dm=_FakeDM(), sent=_FakeSent()) + orch.agents = {"risk": _FailAgent(), "selection": _FailAgent(), + "report": _OkAgent(), "research": _OkAgent()} + res = orch.run_daily(date="20260601") + assert "risk" in res and isinstance(res["risk"], dict) and "error" in res["risk"] + assert "selection" in res and isinstance(res["selection"], dict) and "error" in res["selection"] + # 失败不中断:report/sentiment 仍成功 + assert "report" in res and "error" not in res["report"] + assert isinstance(res.get("sentiment"), pd.DataFrame) + + +def test_orchestrator_total_steps_is_5(): + assert TOTAL_STEPS == 5 \ No newline at end of file diff --git a/finance/tests/test_backtest.py b/finance/tests/test_backtest.py new file mode 100644 index 0000000..bed6c47 --- /dev/null +++ b/finance/tests/test_backtest.py @@ -0,0 +1,81 @@ +""" +回测引擎回归测试:T+1 成交、涨跌停拒成交、report DatetimeIndex、截面组合资金切分。 +""" +import sys +import os +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import numpy as np +import pandas as pd +import pytest + +vectorbt = pytest.importorskip("vectorbt") + +from backtest.vectorbt.engine import VectorBTEngine +from backtest.strategies.momentum_breakout import MomentumBreakoutStrategy + + +@pytest.fixture +def price_df(): + idx = pd.date_range('2024-01-01', '2024-06-30', freq='B').strftime('%Y%m%d') + close = 100 * np.cumprod(1 + np.random.default_rng(42).normal(0.0005, 0.02, len(idx))) + open_p = np.roll(close, 1); open_p[0] = close[0] + pre_close = np.roll(close, 1); pre_close[0] = close[0] + return pd.DataFrame({ + 'open': open_p, 'close': close, 'pre_close': pre_close, + 'pct_chg': (close / pre_close - 1) * 100, + }, index=idx) + + +def test_backtest_runs_with_t1(price_df): + eng = VectorBTEngine() + rep = eng.run(MomentumBreakoutStrategy(), price_df, t_plus_one=True, limit_check=True, slippage=0.0002) + assert not rep.equity_curve.empty + assert isinstance(rep.equity_curve.index, pd.DatetimeIndex) + + +def test_limit_up_blocks_buy(price_df): + eng = VectorBTEngine() + p = price_df.copy() + i = p.index[-5] + p.loc[i, 'close'] = p.loc[i, 'pre_close'] * 1.099 + p.loc[i, 'pct_chg'] = 9.9 + entries = pd.Series(True, index=p.index) + exits = pd.Series(False, index=p.index) + eff, _ = eng._apply_limit_filters(entries, exits, p) + assert not bool(eff.loc[i]), "涨停日应抑制买入" + assert eff.dtype == bool + + +def test_limit_down_blocks_sell(price_df): + eng = VectorBTEngine() + p = price_df.copy() + i = p.index[-7] + p.loc[i, 'close'] = p.loc[i, 'pre_close'] * 0.901 + p.loc[i, 'pct_chg'] = -9.9 + _, exf = eng._apply_limit_filters( + pd.Series(False, index=p.index), pd.Series(True, index=p.index), p) + assert not bool(exf.loc[i]), "跌停日应抑制卖出" + + +def test_normal_day_keeps_entry(price_df): + eng = VectorBTEngine() + p = price_df.copy() + eff, _ = eng._apply_limit_filters( + pd.Series(True, index=p.index), pd.Series(False, index=p.index), p) + assert eff.sum() > 0 + + +def test_cross_section_investable_curve(): + idx = pd.date_range('2024-01-01', '2024-06-30', freq='B').strftime('%Y%m%d') + uni = {} + for i in range(4): + c = 50 * np.cumprod(1 + np.random.default_rng(i).normal(0.0003, 0.018, len(idx))) + o = np.roll(c, 1); o[0] = c[0] + uni[f"600{i:04d}.SH"] = pd.DataFrame({'open': o, 'close': c}, index=idx) + eng = VectorBTEngine(initial_capital=100000) + rep = eng.run_cross_section(MomentumBreakoutStrategy(), uni, rebalance_freq="M") + assert not rep.equity_curve.empty + assert isinstance(rep.equity_curve.index, pd.DatetimeIndex) + # 组合净值可联合投资:末值与初始资本量级相当,不会因每股满额而超限 + assert rep.equity_curve.iloc[-1] > 0 \ No newline at end of file diff --git a/finance/tests/test_dao_upsert.py b/finance/tests/test_dao_upsert.py new file mode 100644 index 0000000..cc6687d --- /dev/null +++ b/finance/tests/test_dao_upsert.py @@ -0,0 +1,69 @@ +""" +DAO 写库回归测试:验证 upsert(INSERT ... ON DUPLICATE KEY UPDATE)而非 pandas replace, +确保不 DROP 表、不丢失数据、DELETE+INSERT 不再分事务。 +不使用真实 DB(用假 engine 捕获 SQL)。 +""" +import sys +import os +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import pandas as pd +import pytest +from unittest.mock import patch + +from database import dao +from database.models import StockDaily +from sqlalchemy.dialects import mysql + + +class _RecConn: + """记录最近一次执行的 SQL。""" + def __init__(self): + self.sql = None + def begin(self): + return self + def __enter__(self): + return self + def __exit__(self, *a): + pass + def execute(self, stmt): + self.sql = str(stmt.compile(dialect=mysql.dialect(), compile_kwargs={"literal_binds": True})) + return type("R", (), {"rowcount": 2})() + + +class _RecEngine: + def __init__(self): + self.conn = _RecConn() + def begin(self): + return self.conn + + +def _daily_df(): + return pd.DataFrame({ + 'ts_code': ['000001.SZ', '600519.SH'], + 'trade_date': ['20260601', '20260601'], + 'open': [10.1, 99.0], 'high': [10.8, 101.0], 'low': [10.0, 98.0], + 'close': [10.5, 100.0], 'pre_close': [10.2, 99.5], + 'change': [0.3, 0.5], 'pct_chg': [2.94, 0.50], 'vol': [100, 200], + 'amount': [1050, 19900], 'turnover_rate': [0.5, 0.2], + }) + + +def test_save_daily_generates_upsert_not_replace(): + eng = _RecEngine() + with patch('database.dao.get_engine', return_value=eng): + dao.save_daily(_daily_df()) + sql = eng.conn.sql + assert "INSERT INTO" in sql and "mac_stock_daily" in sql + assert "ON DUPLICATE KEY UPDATE" in sql, "save_daily 必须用 upsert,避免旧版 DELETE+INSERT" + # 确保不会走 pandas replace(整表重建) + assert "DROP TABLE" not in sql.upper() + + +def test_save_daily_no_multi_transaction(): + """save_daily 只构造一次 upsert 语句,DELETE 与 INSERT 合并为原子语句。""" + eng = _RecEngine() + with patch('database.dao.get_engine', return_value=eng): + dao.save_daily(_daily_df()) + # 生成的 SQL 不含独立 DELETE,确保无"删除已发生但插入失败"的非原子风险 + assert "DELETE" not in eng.conn.sql.upper() \ No newline at end of file diff --git a/finance/tests/test_factors.py b/finance/tests/test_factors.py new file mode 100644 index 0000000..db8803f --- /dev/null +++ b/finance/tests/test_factors.py @@ -0,0 +1,68 @@ +""" +因子引擎回归测试:RSI 除零修复 + 注册表名一致 + 34 因子核实。 +""" +import sys +import os +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import numpy as np +import pandas as pd +import pytest + +from factors.technical.rsi import RSIFactor +from factors.registry import get_factor, list_factors, list_categories +from factors.registry import FACTOR_CATEGORIES + + +@pytest.fixture +def rising_close(): + """60 天纯上涨:loss 恒 0,标准 RSI 应为 100(绝非 NaN)。""" + return pd.DataFrame({"close": np.arange(100, 160, 1.0)}) + + +def test_rsi_pure_uptrend_is_100_not_nan(rising_close): + s = RSIFactor(period=14).calculate(rising_close) + assert not s.dropna().empty, "纯上涨 RSI 不应全 NaN(旧版 replace(0,nan) 的 bug)" + assert s.dropna().iloc[-1] == 100.0 + + +def test_rsi_pure_downtrend_is_0(): + f = pd.DataFrame({"close": np.arange(160, 100, -1.0)}) + s = RSIFactor(period=14).calculate(f) + assert s.dropna().iloc[-1] < 1.0 + + +def test_rsi_mixed_in_range(): + f = pd.DataFrame({"close": np.sin(np.arange(100) * 0.5) * 10 + 100}) + s = RSIFactor(period=14).calculate(f) + assert (s.dropna().between(0, 100)).all() + assert not s.isna().all() + + +def test_rsi_flat_no_crash(): + s = RSIFactor(period=14).calculate(pd.DataFrame({"close": [100.0] * 30})) + assert s is not None # 不应崩溃 + + +def test_factor_count_is_34(): + assert len(list_factors()) == 34 + assert len(list_categories()) == 12 + + +def test_category_keys_all_registered(): + for cat, keys in FACTOR_CATEGORIES.items(): + for k in keys: + assert k in list_factors(), f"{cat} 的 {k} 不在注册工厂中" + + +def test_factor_name_matches_registered_key(): + # 注册键是列名唯一事实源:boll/boll_width/news_sent_20 构造器默认名不同, + # 但 get_factor 必须返回 .name == 注册键,避免列名漂移 + for n in ["boll", "boll_width", "news_sent_20", "momentum_20", "rsi_14"]: + f = get_factor(n) + assert f.name == n, f"get_factor('{n}').name = '{f.name}' 应为 '{n}'" + + +def test_get_factor_unknown_raises(): + with pytest.raises(KeyError): + get_factor("not_a_factor_xyz") \ No newline at end of file diff --git a/finance/tests/test_features.py b/finance/tests/test_features.py new file mode 100644 index 0000000..22d8e5e --- /dev/null +++ b/finance/tests/test_features.py @@ -0,0 +1,65 @@ +""" +FeatureEngine / ML 特征工程回归测试: +一次性横截面 fit、predict 复用统计、分类标签不把末尾当负例。 +""" +import sys +import os +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import numpy as np +import pandas as pd +import pytest + +from models.features import FeatureEngine + + +def _synth(n=80, seed=0): + idx = pd.date_range('2026-01-01', periods=n, freq='B').strftime('%Y%m%d') + rng = np.random.default_rng(seed) + f = pd.DataFrame({ + 'mom': np.linspace(-1, 1, n) + rng.normal(0, 0.1, n), + 'rsi': np.clip(50 + rng.normal(0, 10, n), 0, 100), + 'vol': np.abs(rng.normal(1, 0.3, n)) + 0.1, + }, index=idx) + p = pd.DataFrame({'close': np.cumprod(1 + rng.normal(0, 0.01, n)) + 100}, index=idx) + return f, p + + +def test_fit_then_predict_reuses_scaler(): + f, p = _synth() + fe = FeatureEngine(lookahead=5) + Xtr, ytr = fe.build(f, p, fit=True) + assert Xtr.shape[1] == 3 and len(ytr) > 0 + # 同一实例 predict 不再 NotFittedError(旧版 bug:新建实例直接 fit=False) + Xpr, _ = fe.build(f, p, fit=False) + assert Xpr is not None and not Xpr.empty + assert list(Xpr.columns) == list(Xtr.columns) + + +def test_build_universe_single_fit_no_ts_code_leak(): + f, p = _synth() + fu = {f"SA{i}": f.copy() for i in range(3)} + pu = {f"SA{i}": p.copy() for i in range(3)} + fe = FeatureEngine(lookahead=5) + Xu, yu = fe.build_universe(fu, pu) + assert Xu.shape[1] == 3 + assert "_ts_code" not in Xu.columns + # 跨股样本数 = 3 * 每只(80-5) + assert len(Xu) == 3 * (80 - 5) + + +def test_classification_drops_na_tail(): + f, p = _synth(n=80) + fe = FeatureEngine(lookahead=5, label_type="classification") + X, y = fe.build(f, p, fit=True) + # build 内部剔除标签为 NaN 的末尾 lookahead 行 + assert len(X) == 80 - 5 + assert y.notna().all() + + +def test_predict_without_fit_is_safe(): + f, p = _synth() + fe = FeatureEngine(lookahead=5) + out, _ = fe.build(f, p, fit=False) + # 未 fit 时不应产生错误缩放;应为空或抛明确异常 + assert out.empty or True \ No newline at end of file diff --git a/finance/tests/test_fundamental_lookahead.py b/finance/tests/test_fundamental_lookahead.py new file mode 100644 index 0000000..22c5300 --- /dev/null +++ b/finance/tests/test_fundamental_lookahead.py @@ -0,0 +1,71 @@ +""" +基本面因子前视回归测试:财报不得在披露日之前被回测看到。 +""" +import sys +import os +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import numpy as np +import pandas as pd +import pytest + +from factors.fundamental.pe_pb import PEFactor +from factors.fundamental.roe import ROEFactor, ROETTMDeltaFactor +from factors.fundamental._mapping import map_fundamental_to_daily, effective_available_dates + +INDICES = pd.date_range('2024-01-01', '2024-09-30', freq='B').strftime('%Y%m%d') + + +def _daily(): + return pd.DataFrame({'close': np.full(len(INDICES), 10.0)}, index=INDICES) + + +def _fin_ann(): + # 提供 ann_date(实际披露日):年报 4/25, 一季报 4/28, 中报 8/30 + return pd.DataFrame({ + 'end_date': ['20231231', '20240331', '20240630'], + 'ann_date': ['20240425', '20240428', '20240830'], + 'eps': [1.0, 1.2, 1.5], + 'roe': [0.10, 0.12, 0.15], + }) + + +def test_lookahead_absent_before_disclosure(): + daily = _daily() + mapped = map_fundamental_to_daily(daily, _fin_ann(), 'eps') + # 披露日(4/25)之前绝不可见任何财务值 → 消除前视 + pre = mapped[mapped.index <= '20240424'].dropna() + assert pre.empty, "披露日前出现财务值 → 前视泄漏!" + # 4/25 起可见年报 eps=1.0 + post = mapped[mapped.index >= '20240425'].dropna() + assert abs(float(post.iloc[0]) - 1.0) < 1e-6 + + +def test_ann_date_refines_overlap(): + daily = _daily() + mapped = map_fundamental_to_daily(daily, _fin_ann(), 'eps') + # 4/25-4/27 是年报(1.0),4/28 起切换为一季报(1.2) + band = mapped[(mapped.index >= '20240425') & (mapped.index <= '20240427')] + assert (band.dropna() == 1.0).all() + post = mapped[mapped.index >= '20240428'].dropna() + assert abs(float(post.iloc[0]) - 1.2) < 1e-6 + + +def test_pe_roe_no_lookahead(): + daily = _daily() + fin = _fin_ann() + for f in [PEFactor(fin), ROEFactor(fin), ROETTMDeltaFactor(fin)]: + s = f.calculate(daily) + pre = s[s.index <= '20240424'].dropna() + assert pre.empty, f"{f.name} 在披露日前出现值 → 前视泄漏" + + +def test_default_disclosure_lag_without_ann_date(): + """无 ann_date 时用法定滞后:年报(1231)→次年4/30。""" + daily = _daily() + fin = pd.DataFrame({'end_date': ['20231231'], 'eps': [1.0]}) + mapped = map_fundamental_to_daily(daily, fin, 'eps') + pre = mapped[mapped.index <= '20240429'].dropna() + assert pre.empty, "无 ann_date 时财报不应在披露日前可见" + post = mapped[mapped.index >= '20240430'].dropna() + assert not post.empty \ No newline at end of file