data_manager.py
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
数据管理模块
负责历史日线数据下载/增量更新、当日14:30实时数据获取、实时消息爬取、SQLite数据库管理
"""
import os
import logging
from pathlib import Path
import sqlite3
import pandas as pd
import tushare as ts
import requests
from bs4 import BeautifulSoup
from datetime import datetime, timedelta
from dotenv import load_dotenv
import time
from src.config_manager import config_manager
from src.utils import ensure_dir_exists, retry_on_exception
# 加载环境变量
load_dotenv()
# 配置日志
logger = logging.getLogger(__name__)
class DataManager:
"""数据管理类"""
def __init__(self):
"""初始化数据管理对象"""
self.db_path = config_manager.get('system.db_path')
self.tushare_token = config_manager.get('tushare.token')
self.min_turnover_rate = config_manager.get('selection.min_turnover_rate')
self.max_stock_count = config_manager.get('selection.max_stock_count')
# 初始化Tushare
ts.set_token(self.tushare_token)
self.pro = ts.pro_api()
# 确保数据目录存在
ensure_dir_exists(str(Path(self.db_path).parent))
# 初始化数据库
self.init_database()
def init_database(self):
"""初始化SQLite数据库"""
try:
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
# 创建股票基本信息表
cursor.execute('''
CREATE TABLE IF NOT EXISTS stock_basic (
ts_code TEXT PRIMARY KEY,
name TEXT,
industry TEXT,
market TEXT,
list_date TEXT,
delist_date TEXT,
is_st INTEGER DEFAULT 0
)
''')
# 创建日线数据表
cursor.execute('''
CREATE TABLE IF NOT EXISTS daily_data (
ts_code TEXT,
trade_date TEXT,
open REAL,
high REAL,
low REAL,
close REAL,
pre_close REAL,
change REAL,
pct_chg REAL,
vol REAL,
amount REAL,
PRIMARY KEY (ts_code, trade_date)
)
''')
# 创建实时数据表
cursor.execute('''
CREATE TABLE IF NOT EXISTS realtime_data (
ts_code TEXT,
trade_date TEXT,
trade_time TEXT,
price REAL,
volume REAL,
amount REAL,
bid_volume REAL,
ask_volume REAL,
turnover_rate REAL,
volume_ratio REAL,
PRIMARY KEY (ts_code, trade_date, trade_time)
)
''')
# 创建消息数据表
cursor.execute('''
CREATE TABLE IF NOT EXISTS news_data (
id INTEGER PRIMARY KEY AUTOINCREMENT,
title TEXT,
content TEXT,
source TEXT,
publish_time TEXT,
trade_date TEXT,
industry TEXT,
stock_codes TEXT
)
''')
# 创建索引
cursor.execute('''
CREATE INDEX IF NOT EXISTS idx_daily_data ON daily_data (ts_code, trade_date)
''')
cursor.execute('''
CREATE INDEX IF NOT EXISTS idx_realtime_data ON realtime_data (ts_code, trade_date, trade_time)
''')
cursor.execute('''
CREATE INDEX IF NOT EXISTS idx_news_data ON news_data (trade_date, industry)
''')
conn.commit()
conn.close()
logger.info("数据库初始化完成")
except Exception as e:
logger.error(f"数据库初始化失败:{e}")
@retry_on_exception(max_retries=3, delay=1)
def update_stock_basic(self):
"""更新股票基本信息"""
try:
logger.info("开始更新股票基本信息...")
# 获取股票列表
df_basic = self.pro.stock_basic(exchange='', list_status='L', fields='ts_code,name,industry,market,list_date,delist_date')
# 识别ST股票 - 使用简单的方法避免正则表达式问题
df_basic['is_st'] = df_basic['name'].apply(lambda x: 1 if 'ST' in x else 0)
# 保存到数据库
conn = sqlite3.connect(self.db_path)
df_basic.to_sql('stock_basic', conn, if_exists='replace', index=False)
conn.commit()
conn.close()
logger.info(f"股票基本信息更新完成,共 {len(df_basic)} 只股票")
return df_basic
except Exception as e:
logger.error(f"股票基本信息更新失败:{e}")
return pd.DataFrame()
def get_valid_stocks(self):
"""获取有效股票列表(剔除ST、停牌等)"""
try:
conn = sqlite3.connect(self.db_path)
query = '''
SELECT ts_code, name, industry
FROM stock_basic
WHERE is_st = 0
LIMIT ?
'''
df_valid = pd.read_sql_query(query, conn, params=(self.max_stock_count,))
conn.close()
logger.info(f"有效股票数量:{len(df_valid)}")
# 输出前5个股票的代码
if len(df_valid) > 0:
logger.info(f"前5个股票代码:{df_valid['ts_code'].head().tolist()}")
return df_valid
except Exception as e:
logger.error(f"获取有效股票列表失败:{e}")
return pd.DataFrame()
@retry_on_exception(max_retries=3, delay=1)
def _get_trading_dates(self, start_date, end_date):
"""获取指定日期范围内的交易日"""
try:
logger.info(f"获取 {start_date} 到 {end_date} 的交易日...")
# 使用Tushare获取交易日数据
df_trade_cal = self.pro.trade_cal(
start_date=start_date,
end_date=end_date,
is_open=1
)
if not df_trade_cal.empty:
trade_dates = df_trade_cal['cal_date'].tolist()
logger.info(f"成功获取 {len(trade_dates)} 个交易日")
return trade_dates
else:
logger.warning("未获取到交易日数据")
return []
except Exception as e:
logger.error(f"获取交易日数据失败:{e}")
return []
@retry_on_exception(max_retries=3, delay=1)
def update_history_data(self, start_date=None, end_date=None):
"""更新历史日线数据"""
try:
logger.info("开始更新历史日线数据...")
# 获取有效股票列表
df_valid = self.get_valid_stocks()
if df_valid.empty:
logger.warning("无有效股票数据,更新终止")
return
# 确定日期范围
if not start_date:
start_date = '20230101' # 从2023年1月开始更新
if not end_date:
# 如果是当天,且当天还没有交易数据,则结束日期为前一天
today = datetime.now().strftime('%Y%m%d')
if today == datetime.now().strftime('%Y%m%d'):
# 检查当天是否有交易数据
df_today = self.pro.daily(trade_date=today)
if df_today.empty:
end_date = (datetime.now() - timedelta(days=1)).strftime('%Y%m%d')
else:
end_date = today
else:
end_date = today
logger.info(f"更新数据日期范围:{start_date} 到 {end_date}")
# 获取所有交易日
trade_dates = self._get_trading_dates(start_date, end_date)
if not trade_dates:
logger.warning("指定日期范围内无交易日")
return
# 检查数据库中已有的交易日,避免重复更新
conn = sqlite3.connect(self.db_path, timeout=10)
cursor = conn.cursor()
cursor.execute('''
SELECT DISTINCT trade_date FROM daily_data
''')
existing_trade_dates = set([row[0] for row in cursor.fetchall()])
# 过滤掉已有的交易日
trade_dates = [date for date in trade_dates if date not in existing_trade_dates]
logger.info(f"共 {len(trade_dates)} 个交易日需要更新")
# 按日期更新数据,每次更新一个交易日的数据
for i, trade_date in enumerate(trade_dates):
logger.info(f"正在更新第 {i+1}/{len(trade_dates)} 个交易日的数据:{trade_date}")
try:
# 获取当日所有股票的日线数据
df_daily = self.pro.daily(trade_date=trade_date)
if not df_daily.empty:
# 筛选有效股票的数据
df_daily_valid = df_daily[df_daily['ts_code'].isin(df_valid['ts_code'])]
if not df_daily_valid.empty:
# 使用临时表来避免UNIQUE约束失败
df_daily_valid.to_sql('temp_daily', conn, if_exists='replace', index=False)
# 先删除该交易日的所有数据,再插入新数据
cursor = conn.cursor()
cursor.execute('''
DELETE FROM daily_data WHERE trade_date = ?
''', (trade_date,))
cursor.execute('''
INSERT INTO daily_data (ts_code, trade_date, open, high, low, close, pre_close, change, pct_chg, vol, amount)
SELECT ts_code, trade_date, open, high, low, close, pre_close, change, pct_chg, vol, amount FROM temp_daily
''')
conn.commit()
logger.info(f"成功保存 {len(df_daily_valid)} 条数据到数据库")
else:
logger.warning(f"交易日 {trade_date} 无有效股票数据")
else:
logger.warning(f"交易日 {trade_date} 无数据")
except Exception as e:
logger.error(f"获取交易日 {trade_date} 数据失败:{e}")
continue
# 等待一段时间,避免API调用频率过高
time.sleep(0.75)
# 删除重复的数据
logger.info("删除重复的数据...")
cursor.execute('''
DELETE FROM daily_data
WHERE rowid NOT IN (
SELECT MIN(rowid)
FROM daily_data
GROUP BY ts_code, trade_date
)
''')
conn.commit()
# 删除非2023年1月之后的交易日数据
logger.info("删除非2023年1月之后的交易日数据...")
cursor.execute('''
DELETE FROM daily_data
WHERE trade_date < '20230101'
''')
conn.commit()
# 检查是否所有交易日的数据都更新成功
cursor.execute('''
SELECT DISTINCT trade_date FROM daily_data
''')
updated_trade_dates = set([row[0] for row in cursor.fetchall()])
missing_trade_dates = [date for date in trade_dates if date not in updated_trade_dates]
if missing_trade_dates:
logger.warning(f"以下交易日的数据未更新成功:{missing_trade_dates}")
conn.close()
logger.info("历史日线数据更新完成")
except Exception as e:
logger.error(f"更新历史日线数据失败:{e}")
@retry_on_exception(max_retries=3, delay=1)
def get_real_time_data(self, trade_date=None, trade_time=None):
"""获取当日14:30实时数据"""
try:
logger.info("开始获取实时数据...")
if not trade_date:
trade_date = datetime.now().strftime('%Y%m%d')
if not trade_time:
trade_time = '14:30:00'
# 直接使用Tushare的daily方法获取所有股票的数据
logger.info("直接使用Tushare API获取所有股票的日线数据")
df_realtime = self.pro.daily(trade_date=trade_date)
logger.info(f"获取到日线数据数量:{len(df_realtime)}")
if not df_realtime.empty:
df_realtime['trade_time'] = trade_time
# 计算实时量比和换手率
df_realtime['volume_ratio'] = self._calculate_volume_ratio(df_realtime)
df_realtime['turnover_rate'] = self._calculate_turnover_rate(df_realtime)
logger.info(f"平均成交量:{df_realtime['vol'].mean()}")
logger.info(f"平均换手率:{df_realtime['turnover_rate'].mean()}")
# 映射列名到实时数据表的结构
df_realtime = df_realtime.rename(columns={
'close': 'price',
'vol': 'volume'
})
# 选择实时数据表的列
df_realtime = df_realtime[['ts_code', 'trade_date', 'trade_time', 'price', 'volume', 'amount', 'turnover_rate', 'volume_ratio']]
logger.info(f"获取到实时数据数量:{len(df_realtime)}")
# 保存到数据库,使用if_exists='replace'避免重复
conn = sqlite3.connect(self.db_path)
df_realtime.to_sql('realtime_data', conn, if_exists='replace', index=False)
conn.commit()
conn.close()
logger.info(f"实时数据获取完成,共 {len(df_realtime)} 条有效记录")
else:
logger.warning("未获取到实时数据")
return df_realtime
except Exception as e:
logger.error(f"实时数据获取失败:{e}")
return pd.DataFrame()
def _calculate_volume_ratio(self, df):
"""计算量比"""
# 简单模拟量比计算
return (df['vol'] / df['vol'].mean()).clip(0.1, 10)
def _calculate_turnover_rate(self, df):
"""计算换手率"""
# 简单模拟换手率计算
return (df['vol'] / 1000000).clip(0.01, 20)
@retry_on_exception(max_retries=3, delay=1)
def get_news_data(self, trade_date=None):
"""获取实时消息数据"""
try:
logger.info("开始获取消息数据...")
if not trade_date:
trade_date = datetime.now().strftime('%Y%m%d')
# 爬取财联社新闻
news_list = []
# 财联社新闻接口
cailian_url = "https://www.cls.cn/api/sw?app=CailianZb&os=web&sv=7.7.5"
try:
response = requests.get(cailian_url, timeout=10)
if response.status_code == 200:
data = response.json()
if data.get('data'):
for item in data['data']:
news = {
'title': item.get('title', ''),
'content': item.get('content', ''),
'source': '财联社',
'publish_time': item.get('time', ''),
'trade_date': trade_date,
'industry': '',
'stock_codes': ''
}
news_list.append(news)
logger.info(f"财联社新闻获取完成,共 {len(news_list)} 条")
except Exception as e:
logger.error(f"财联社新闻获取失败:{e}")
# 爬取东方财富新闻
eastmoney_url = "https://push2.eastmoney.com/api/qt/clist/get"
params = {
'pn': 1,
'pz': 50,
'po': 1,
'np': 1,
'ut': 'bd1d9ddb04089700cf9c27f6f7426281',
'fltt': 2,
'invt': 2,
'fid': 'f3',
'fs': 'm:0 t:6 f:!2,m:0 t:80 f:!2,m:1 t:2 f:!2,m:1 t:23 f:!2,m:0 t:81 s:2048',
'fields': 'f1,f2,f3,f4,f5,f6,f7,f8,f9,f10,f12,f13,f14,f15,f16,f17,f18,f20,f21,f23,f24,f25,f22,f11,f62,f128,f136,f115,f152',
'_': '1640000000000'
}
try:
response = requests.get(eastmoney_url, params=params, timeout=10)
if response.status_code == 200:
data = response.json()
if data.get('data') and data['data'].get('diff'):
for item in data['data']['diff']:
news = {
'title': item.get('f14', ''),
'content': '',
'source': '东方财富',
'publish_time': item.get('f13', ''),
'trade_date': trade_date,
'industry': '',
'stock_codes': ''
}
news_list.append(news)
logger.info(f"东方财富新闻获取完成,共 {len(news_list)} 条")
except Exception as e:
logger.error(f"东方财富新闻获取失败:{e}")
# 保存到数据库
if news_list:
df_news = pd.DataFrame(news_list)
conn = sqlite3.connect(self.db_path)
df_news.to_sql('news_data', conn, if_exists='append', index=False)
conn.commit()
conn.close()
logger.info(f"消息数据保存完成,共 {len(news_list)} 条")
else:
logger.warning("未获取到消息数据")
return pd.DataFrame(news_list)
except Exception as e:
logger.error(f"消息数据获取失败:{e}")
return pd.DataFrame()
def get_stock_data_for_model(self, ts_code, start_date=None, end_date=None):
"""获取单只股票的历史双数据(日线+实时)"""
try:
if not end_date:
end_date = datetime.now().strftime('%Y%m%d')
if not start_date:
start_date = (datetime.now() - timedelta(days=120)).strftime('%Y%m%d')
conn = sqlite3.connect(self.db_path)
# 获取日线数据
daily_query = '''
SELECT trade_date, open, high, low, close, vol, amount, pct_chg
FROM daily_data
WHERE ts_code = ? AND trade_date >= ? AND trade_date <= ?
ORDER BY trade_date
'''
df_daily = pd.read_sql_query(daily_query, conn, params=(ts_code, start_date, end_date))
# 获取实时数据
realtime_query = '''
SELECT trade_date, trade_time, price, volume, amount, turnover_rate, volume_ratio
FROM realtime_data
WHERE ts_code = ? AND trade_date >= ? AND trade_date <= ?
ORDER BY trade_date, trade_time
'''
df_realtime = pd.read_sql_query(realtime_query, conn, params=(ts_code, start_date, end_date))
conn.close()
# 合并数据
if not df_daily.empty and not df_realtime.empty:
df_combined = pd.merge(
df_daily,
df_realtime,
on='trade_date',
how='inner'
)
return df_combined
elif not df_daily.empty:
# 如果只有日线数据,也返回日线数据(用于回测)
return df_daily
else:
return pd.DataFrame()
except Exception as e:
logger.error(f"获取股票 {ts_code} 数据失败:{e}")
return pd.DataFrame()
def get_all_stock_data_for_backtest(self, start_date=None, end_date=None):
"""获取所有股票的历史双数据(用于回测)"""
try:
if not end_date:
end_date = datetime.now().strftime('%Y%m%d')
if not start_date:
start_date = (datetime.now() - timedelta(days=365)).strftime('%Y%m%d')
conn = sqlite3.connect(self.db_path)
# 获取所有有效股票
df_valid = self.get_valid_stocks()
all_stock_data = {}
for ts_code in df_valid['ts_code']:
df_stock = self.get_stock_data_for_model(ts_code, start_date, end_date)
if not df_stock.empty:
all_stock_data[ts_code] = df_stock
conn.close()
logger.info(f"回测数据获取完成,共 {len(all_stock_data)} 只股票有有效数据")
return all_stock_data
except Exception as e:
logger.error(f"回测数据获取失败:{e}")
return {}
def test_connection(self):
"""测试连接"""
try:
# 测试Tushare连接
test_df = self.pro.stock_basic(exchange='', list_status='L', fields='ts_code', limit=10)
logger.info(f"Tushare连接测试通过,获取到 {len(test_df)} 条股票数据")
# 测试数据库连接
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM stock_basic")
count = cursor.fetchone()[0]
logger.info(f"数据库连接测试通过,股票基本信息表有 {count} 条记录")
conn.close()
return True
except Exception as e:
logger.error(f"连接测试失败:{e}")
return False
if __name__ == "__main__":
# 配置日志
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
# 测试数据管理模块
data_manager = DataManager()
data_manager.test_connection()
data_manager.update_stock_basic()
data_manager.update_history_data()
data_manager.get_real_time_data()
data_manager.get_news_data()