SHA256
275 lines
10 KiB
Python
275 lines
10 KiB
Python
"""数据库模型定义与连接管理
|
||
|
||
表结构概览:
|
||
- stock_info: 股票基本信息(含上市日期,用于跳过未上市股票)
|
||
- stock_daily: 日线行情(多源抓取:BaoStock/新浪/腾讯)
|
||
- stock_financial_income/balance/cashflow: 三大财务报表(JSON存储)
|
||
- stock_money_flow: 个股资金流向
|
||
- stock_dragon_tiger: 龙虎榜
|
||
- stock_dividend: 分红送转
|
||
- stock_intraday: 1分钟分时行情
|
||
- trading_day: 交易日历(用于判断数据完整性)
|
||
- stock_no_data: 无数据/停牌记录(避免重复抓取)
|
||
"""
|
||
|
||
from sqlalchemy import (
|
||
Column, String, Date, DateTime, Float, Integer, Text,
|
||
UniqueConstraint, Index, create_engine, MetaData, func, text,
|
||
)
|
||
from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker
|
||
from sqlalchemy.dialects.mysql import insert as mysql_insert
|
||
|
||
from src.config import get_mysql_url
|
||
|
||
|
||
class Base(DeclarativeBase):
|
||
pass
|
||
|
||
|
||
# ── 股票基本信息 ──
|
||
class StockInfo(Base):
|
||
__tablename__ = "stock_info"
|
||
|
||
code = Column(String(10), primary_key=True, comment="股票代码")
|
||
name = Column(String(50), comment="股票名称")
|
||
# 上市日期用于在抓取历史数据时跳过当时尚未上市的股票
|
||
ipo_date = Column(Date, comment="上市日期")
|
||
|
||
|
||
# ── 日线行情 ──
|
||
class StockDaily(Base):
|
||
__tablename__ = "stock_daily"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "date", name="uq_daily_code_date"),
|
||
Index("ix_daily_date", "date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
# 多个数据源(BaoStock/新浪/腾讯)写入同一张表,通过 upsert 去重
|
||
date = Column(Date, nullable=False, comment="交易日期")
|
||
open = Column(Float, comment="开盘价")
|
||
close = Column(Float, comment="收盘价")
|
||
high = Column(Float, comment="最高价")
|
||
low = Column(Float, comment="最低价")
|
||
volume = Column(Float, comment="成交量")
|
||
turnover = Column(Float, comment="成交额")
|
||
amplitude = Column(Float, comment="振幅%")
|
||
pct_change = Column(Float, comment="涨跌幅%")
|
||
change = Column(Float, comment="涨跌额")
|
||
turnover_rate = Column(Float, comment="换手率%")
|
||
|
||
|
||
# ── 利润表 ──
|
||
class FinancialIncome(Base):
|
||
__tablename__ = "stock_financial_income"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "report_date", name="uq_income_code_date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
report_date = Column(String(20), nullable=False, comment="报告期")
|
||
data = Column(Text, comment="JSON格式利润表数据")
|
||
|
||
|
||
# ── 资产负债表 ──
|
||
class FinancialBalance(Base):
|
||
__tablename__ = "stock_financial_balance"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "report_date", name="uq_balance_code_date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
report_date = Column(String(20), nullable=False, comment="报告期")
|
||
data = Column(Text, comment="JSON格式资产负债表数据")
|
||
|
||
|
||
# ── 现金流量表 ──
|
||
class FinancialCashflow(Base):
|
||
__tablename__ = "stock_financial_cashflow"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "report_date", name="uq_cashflow_code_date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
report_date = Column(String(20), nullable=False, comment="报告期")
|
||
data = Column(Text, comment="JSON格式现金流量表数据")
|
||
|
||
|
||
# ── 资金流向 ──
|
||
class StockMoneyFlow(Base):
|
||
__tablename__ = "stock_money_flow"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "date", name="uq_moneyflow_code_date"),
|
||
Index("ix_moneyflow_date", "date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
date = Column(Date, nullable=False, comment="日期")
|
||
close = Column(Float, comment="收盘价")
|
||
pct_change = Column(Float, comment="涨跌幅%")
|
||
main_net_inflow = Column(Float, comment="主力净流入-净额")
|
||
main_net_pct = Column(Float, comment="主力净流入-净占比")
|
||
huge_net_inflow = Column(Float, comment="超大盘净流入-净额")
|
||
huge_net_pct = Column(Float, comment="超大盘净流入-净占比")
|
||
big_net_inflow = Column(Float, comment="大盘净流入-净额")
|
||
big_net_pct = Column(Float, comment="大盘净流入-净占比")
|
||
mid_net_inflow = Column(Float, comment="中盘净流入-净额")
|
||
mid_net_pct = Column(Float, comment="中盘净流入-净占比")
|
||
small_net_inflow = Column(Float, comment="小盘净流入-净额")
|
||
small_net_pct = Column(Float, comment="小盘净流入-净占比")
|
||
|
||
|
||
# ── 龙虎榜 ──
|
||
class StockDragonTiger(Base):
|
||
__tablename__ = "stock_dragon_tiger"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "date", name="uq_lhb_code_date"),
|
||
Index("ix_lhb_date", "date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
name = Column(String(50), comment="股票名称")
|
||
date = Column(Date, nullable=False, comment="上榜日期")
|
||
close = Column(Float, comment="收盘价")
|
||
pct_change = Column(Float, comment="涨跌幅%")
|
||
reason = Column(String(200), comment="上榜原因")
|
||
buy_amount = Column(Float, comment="买入额")
|
||
sell_amount = Column(Float, comment="卖出额")
|
||
net_amount = Column(Float, comment="净额")
|
||
|
||
|
||
# ── 分红送转 ──
|
||
class StockDividend(Base):
|
||
__tablename__ = "stock_dividend"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "report_date", name="uq_dividend_code_date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
name = Column(String(50), comment="股票名称")
|
||
report_date = Column(String(20), nullable=False, comment="报告期")
|
||
dividend_date = Column(Date, comment="除权除息日")
|
||
bonus_ratio = Column(Float, comment="每10股送转比例")
|
||
cash_div = Column(Float, comment="每10股派息")
|
||
convert_ratio = Column(Float, comment="每10股转增比例")
|
||
ex_right_date = Column(Date, comment="除权日")
|
||
dividend_yield = Column(Float, comment="股息率%")
|
||
|
||
|
||
# ── 交易日历 ──
|
||
class TradingDay(Base):
|
||
__tablename__ = "trading_day"
|
||
__table_args__ = (
|
||
UniqueConstraint("date", name="uq_trading_day_date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
date = Column(Date, nullable=False, comment="交易日期")
|
||
|
||
|
||
# ── 无数据/停牌记录(按天粒度) ──
|
||
class StockNoData(Base):
|
||
__tablename__ = "stock_no_data"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "date", name="uq_nodata_code_date"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
date = Column(Date, nullable=False, comment="停牌/无数据日期")
|
||
created_at = Column(DateTime, server_default=func.now(), comment="记录时间")
|
||
|
||
|
||
# ── 分时行情(1分钟线) ──
|
||
class StockIntraday(Base):
|
||
__tablename__ = "stock_intraday"
|
||
__table_args__ = (
|
||
UniqueConstraint("code", "datetime", name="uq_intraday_code_dt"),
|
||
Index("ix_intraday_date", "datetime"),
|
||
)
|
||
|
||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||
code = Column(String(10), nullable=False, comment="股票代码")
|
||
datetime = Column(DateTime, nullable=False, comment="时间")
|
||
open = Column(Float, comment="开盘价")
|
||
high = Column(Float, comment="最高价")
|
||
low = Column(Float, comment="最低价")
|
||
close = Column(Float, comment="收盘价")
|
||
volume = Column(Float, comment="成交量")
|
||
amount = Column(Float, comment="成交额")
|
||
|
||
|
||
# ── 数据库连接管理 ──
|
||
_engine = None
|
||
_SessionFactory = None
|
||
|
||
|
||
def get_engine():
|
||
global _engine
|
||
if _engine is None:
|
||
_engine = create_engine(get_mysql_url(), pool_size=5, pool_recycle=3600)
|
||
return _engine
|
||
|
||
|
||
def get_session() -> Session:
|
||
global _SessionFactory
|
||
if _SessionFactory is None:
|
||
_SessionFactory = sessionmaker(bind=get_engine())
|
||
return _SessionFactory()
|
||
|
||
|
||
def init_db():
|
||
engine = get_engine()
|
||
# 自动迁移:旧版 stock_no_data 使用 date_range 列,新版改为 date 列
|
||
# 检测到旧表结构时先删除,由 create_all 重建
|
||
with engine.connect() as conn:
|
||
result = conn.execute(text("SHOW COLUMNS FROM stock_no_data LIKE 'date_range'"))
|
||
if result.fetchone():
|
||
conn.execute(text("DROP TABLE stock_no_data"))
|
||
conn.commit()
|
||
print(" stock_no_data 表结构已升级(date_range → date)")
|
||
Base.metadata.create_all(engine)
|
||
print("数据库表初始化完成")
|
||
|
||
|
||
def batch_upsert(model_cls: type[Base], rows: list[dict], index_columns: list[str]):
|
||
"""MySQL批量upsert:INSERT ON DUPLICATE KEY UPDATE
|
||
|
||
index_columns: 用于判断重复的唯一键列名(如 ["code", "date"]),
|
||
这些列在冲突时不更新,其余列使用新值覆盖。
|
||
"""
|
||
if not rows:
|
||
return
|
||
session = get_session()
|
||
try:
|
||
stmt = mysql_insert(model_cls).values(rows)
|
||
# 只更新输入数据中实际包含的列,排除唯一键列、自增主键和 server_default 列
|
||
input_keys = set(rows[0].keys())
|
||
update_dict = {
|
||
col.name: stmt.inserted[col.name]
|
||
for col in model_cls.__table__.columns
|
||
if col.name in input_keys
|
||
and col.name not in index_columns
|
||
and not col.primary_key
|
||
and not col.server_default
|
||
}
|
||
if update_dict:
|
||
stmt = stmt.on_duplicate_key_update(**update_dict)
|
||
else:
|
||
# 无可更新列时(如 TradingDay 只有唯一键列),用唯一键本身做 no-op 更新
|
||
stmt = stmt.on_duplicate_key_update(**{index_columns[0]: stmt.inserted[index_columns[0]]})
|
||
session.execute(stmt)
|
||
session.commit()
|
||
except Exception as e:
|
||
session.rollback()
|
||
raise e
|
||
finally:
|
||
session.close()
|