""" 数据访问对象。 提供 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})