7 Sprints 全部完成: Sprint 0: 基础设施 (DataManager + MariaDB) Sprint 1: 因子引擎 (34因子/12分类) Sprint 2: VectorBT 回测 (5策略+截面) Sprint 3: Optuna 优化 (+Walk-Forward) Sprint 4: ML 模型 (LightGBM+CatBoost) Sprint 5: Qwen 情绪因子 (三源新闻+日期对齐) Sprint 6: Agent 系统 (4Agent+日报.md/.html) 生产加固 (15项): Tushare双源fallback, SSH自动恢复, pool_pre_ping, save_daily先删后插, load_dotenv绝对路径, 日报5d/20d修复, RiskAgent改上证指数, 昨日对比+数据截止, mac_report utf8mb4, CLAUDE-*.md 9条已知Bug, demo全参数化, djapi数据源归一化, indexDatas API修正 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
134 lines
4.9 KiB
Python
134 lines
4.9 KiB
Python
"""
|
|
数据访问对象。
|
|
|
|
提供 DataFrame 级别的读写操作,屏蔽底层 ORM/SQL 细节。
|
|
"""
|
|
|
|
import pandas as pd
|
|
from sqlalchemy import text
|
|
|
|
from database.connection import get_engine
|
|
from database.models import StockBasic, StockDaily, StockFinancial, Report
|
|
|
|
# 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 _df_to_db(df: pd.DataFrame, model_class, replace: bool = False) -> int:
|
|
"""将 DataFrame 写入对应表,返回写入行数。"""
|
|
if df.empty:
|
|
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
|
|
|
|
|
|
# ── StockBasic ─────────────────────────────────────────────
|
|
|
|
def save_stock_list(df: pd.DataFrame) -> int:
|
|
"""保存股票列表(replace 模式)。"""
|
|
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)
|
|
|
|
|
|
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:
|
|
"""批量写入日线数据。先删旧再插新,避免主键冲突。"""
|
|
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()
|
|
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)
|
|
|
|
|
|
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:
|
|
"""批量写入财务数据(replace 模式:同报告期覆盖更新)。"""
|
|
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)
|
|
|
|
|
|
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})
|