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