- 新增 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 扩充配置项
166 lines
6.1 KiB
Python
166 lines
6.1 KiB
Python
"""
|
||
数据访问对象。
|
||
|
||
提供 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",
|
||
"pre_close", "change", "pct_chg", "vol", "amount", "turnover_rate",
|
||
]
|
||
_FINA_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",
|
||
]
|
||
|
||
|
||
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
|
||
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:
|
||
"""
|
||
保存股票列表(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()
|
||
if df.empty:
|
||
return 0
|
||
engine = get_engine()
|
||
return _upsert_df(engine, StockBasic, df)
|
||
|
||
|
||
def query_stock_list() -> pd.DataFrame:
|
||
"""查询全部股票列表。"""
|
||
engine = get_engine()
|
||
return pd.read_sql(f"SELECT * FROM {StockBasic.__tablename__}", con=engine).set_index("ts_code")
|
||
|
||
|
||
# ── 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",
|
||
]
|
||
df = df[[c for c in cols if c in df.columns]].copy()
|
||
if df.empty:
|
||
return 0
|
||
engine = get_engine()
|
||
return _upsert_df(engine, StockDaily, df)
|
||
|
||
|
||
def query_daily(ts_code: str, start: str | None = None, end: str | None = None) -> pd.DataFrame:
|
||
"""按股票代码和日期范围查询日线。"""
|
||
engine = get_engine()
|
||
table = StockDaily.__tablename__
|
||
sql = f"SELECT * FROM {table} WHERE ts_code = :ts_code"
|
||
params = {"ts_code": ts_code}
|
||
if start:
|
||
sql += " AND trade_date >= :start"
|
||
params["start"] = start
|
||
if end:
|
||
sql += " AND trade_date <= :end"
|
||
params["end"] = end
|
||
sql += " ORDER BY trade_date ASC"
|
||
df = pd.read_sql(text(sql), con=engine, params=params)
|
||
if not df.empty:
|
||
df["trade_date"] = df["trade_date"].astype(str)
|
||
return df
|
||
|
||
|
||
def get_latest_trade_date(ts_code: str) -> str | None:
|
||
"""获取某股票在数据库中的最新交易日。"""
|
||
engine = get_engine()
|
||
table = StockDaily.__tablename__
|
||
sql = f"SELECT MAX(trade_date) FROM {table} WHERE ts_code = :ts_code"
|
||
with engine.connect() as conn:
|
||
result = conn.execute(text(sql), {"ts_code": ts_code}).scalar()
|
||
return result
|
||
|
||
|
||
# ── StockFinancial ─────────────────────────────────────────
|
||
|
||
def save_financial(df: pd.DataFrame) -> int:
|
||
"""
|
||
批量写入财务数据(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()
|
||
if df.empty:
|
||
return 0
|
||
engine = get_engine()
|
||
return _upsert_df(engine, StockFinancial, df)
|
||
|
||
|
||
def query_financial(ts_code: str) -> pd.DataFrame:
|
||
"""查询某股票全部财务数据。"""
|
||
engine = get_engine()
|
||
table = StockFinancial.__tablename__
|
||
sql = f"SELECT * FROM {table} WHERE ts_code = :ts_code ORDER BY end_date DESC"
|
||
return pd.read_sql(text(sql), con=engine, params={"ts_code": ts_code})
|