feat: 量化引擎加固 — 新增测试 + 数据/因子/回测层优化

- 新增 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 扩充配置项
This commit is contained in:
Simon
2026-08-31 14:01:06 +08:00
parent 6acf938caf
commit 73d191b43a
28 changed files with 1418 additions and 373 deletions
+66 -34
View File
@@ -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: