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:
+66
-34
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user