Files
myquant/finance/database/dao.py
T
Simon 73d191b43a 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 扩充配置项
2026-08-31 14:01:06 +08:00

166 lines
6.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
数据访问对象。
提供 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})