data_manager.py
"""
QTrading 数据管理模块
"""
import pandas as pd
from typing import Optional, List
from config.config import config
from data.tushare_api import TushareDataFetcher
from data.database import DatabaseManager
from utils.logger import get_logger
from utils.common import validate_code, format_date, today_str
logger = get_logger(__name__)
class DataManager:
"""数据管理器"""
def __init__(self, tushare_token: str = None):
"""
初始化数据管理器
Args:
tushare_token: Tushare 接口令牌(可选)
"""
self.tushare = TushareDataFetcher(tushare_token)
self.db = DatabaseManager()
def format_ts_code(self, code: str) -> str:
"""
格式化股票代码为 Tushare 格式(6位代码 + .SZ/.SH)
Args:
code: 股票代码(如 '000001' 或 '000001.SZ')
Returns:
str: 格式化后的代码
"""
if '.' in code:
return code
if validate_code(code):
if code.startswith('6'):
return f"{code}.SH"
else:
return f"{code}.SZ"
raise ValueError(f"无效的股票代码: {code}")
def get_kline_data(self, code: str, start_date: str = None, end_date: str = None) -> Optional[pd.DataFrame]:
"""
获取股票 K 线数据(支持增量更新本地缓存)
Args:
code: 股票代码(如 '000001' 或 '000001.SZ')
start_date: 开始日期(YYYYMMDD,可选)
end_date: 结束日期(YYYYMMDD,可选)
Returns:
pd.DataFrame: K 线数据(含 ts_code, trade_date, open, high, low, close, vol, amount)
"""
ts_code = self.format_ts_code(code)
end_date = end_date or today_str()
# 检查本地缓存
last_date = self.db.get_last_trade_date(ts_code)
if last_date:
logger.info(f"本地缓存最后日期: {last_date}")
if last_date >= end_date:
logger.info(f"本地缓存已包含最新数据,直接返回")
return self.db.get_daily_kline(ts_code, start_date, end_date)
# 从 Tushare 获取数据
fetch_start_date = last_date if last_date else (start_date or config.stock_pool.start_date)
kline_data = self.tushare.get_daily_kline(ts_code, fetch_start_date, end_date)
adj_data = self.tushare.get_adj_factor(ts_code, fetch_start_date, end_date)
if kline_data is not None and not kline_data.empty:
# 保存到本地数据库
self.db.save_daily_kline(kline_data)
if adj_data is not None and not adj_data.empty:
self.db.save_adj_factor(adj_data)
# 从数据库获取完整数据
return self.db.get_daily_kline(ts_code, start_date, end_date)
def get_adj_factor_data(self, code: str, start_date: str = None, end_date: str = None) -> Optional[pd.DataFrame]:
"""
获取复权因子数据
Args:
code: 股票代码(如 '000001' 或 '000001.SZ')
start_date: 开始日期(YYYYMMDD,可选)
end_date: 结束日期(YYYYMMDD,可选)
Returns:
pd.DataFrame: 复权因子数据(含 ts_code, trade_date, adj_factor)
"""
ts_code = self.format_ts_code(code)
end_date = end_date or today_str()
# 先检查本地缓存
last_date = self.db.get_last_trade_date(ts_code)
if last_date and last_date >= end_date:
logger.info(f"本地缓存已包含最新复权因子数据")
return self.db.get_adj_factor(ts_code, start_date, end_date)
# 从 Tushare 获取数据
adj_data = self.tushare.get_adj_factor(ts_code, start_date or config.stock_pool.start_date, end_date)
if adj_data is not None and not adj_data.empty:
self.db.save_adj_factor(adj_data)
return self.db.get_adj_factor(ts_code, start_date, end_date)
def calculate_adjusted_data(self, code: str, adjust_method: str = "qfq", start_date: str = None, end_date: str = None) -> Optional[pd.DataFrame]:
"""
计算复权后的 K 线数据
Args:
code: 股票代码(如 '000001' 或 '000001.SZ')
adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
start_date: 开始日期(YYYYMMDD,可选)
end_date: 结束日期(YYYYMMDD,可选)
Returns:
pd.DataFrame: 复权后的 K 线数据
"""
kline_data = self.get_kline_data(code, start_date, end_date)
if kline_data is None or kline_data.empty:
logger.warning(f"未找到股票 {code} 的数据")
return None
if adjust_method == "none":
logger.info(f"使用不复权数据")
return kline_data
adj_data = self.get_adj_factor_data(code, start_date, end_date)
if adj_data is None or adj_data.empty:
logger.warning(f"未找到股票 {code} 的复权因子,返回不复权数据")
return kline_data
# 合并数据
data = pd.merge(kline_data, adj_data, on=['ts_code', 'trade_date'], how='inner')
if adjust_method == "qfq":
return self._calculate_qfq(data)
elif adjust_method == "hfq":
return self._calculate_hfq(data)
else:
logger.warning(f"未知的复权方式: {adjust_method},返回不复权数据")
return kline_data
def _calculate_qfq(self, data: pd.DataFrame) -> pd.DataFrame:
"""计算前复权数据"""
try:
# 取最后一个交易日的复权因子作为基准
last_adj_factor = data['adj_factor'].iloc[-1]
data['adj_factor_qfq'] = data['adj_factor'] / last_adj_factor
data['open'] = data['open'] * data['adj_factor_qfq']
data['high'] = data['high'] * data['adj_factor_qfq']
data['low'] = data['low'] * data['adj_factor_qfq']
data['close'] = data['close'] * data['adj_factor_qfq']
logger.info(f"前复权数据计算完成")
return data
except Exception as e:
logger.error(f"前复权计算失败: {e}")
return data
def _calculate_hfq(self, data: pd.DataFrame) -> pd.DataFrame:
"""计算后复权数据"""
try:
# 取第一个交易日的复权因子作为基准
first_adj_factor = data['adj_factor'].iloc[0]
data['adj_factor_hfq'] = data['adj_factor'] / first_adj_factor
data['open'] = data['open'] * data['adj_factor_hfq']
data['high'] = data['high'] * data['adj_factor_hfq']
data['low'] = data['low'] * data['adj_factor_hfq']
data['close'] = data['close'] * data['adj_factor_hfq']
logger.info(f"后复权数据计算完成")
return data
except Exception as e:
logger.error(f"后复权计算失败: {e}")
return data
def update_all_stocks(self, stock_list: List[str] = None) -> None:
"""
更新股票池数据
Args:
stock_list: 股票代码列表(可选,默认从配置文件读取)
"""
if stock_list is None:
stock_list = config.stock_pool.stock_list
logger.info(f"开始更新 {len(stock_list)} 只股票数据")
success_count = 0
fail_count = 0
for i, code in enumerate(stock_list, 1):
logger.info(f"正在更新 {i}/{len(stock_list)}: {code}")
try:
self.get_kline_data(code)
success_count += 1
except Exception as e:
logger.error(f"更新股票 {code} 失败: {e}")
fail_count += 1
logger.info(f"更新完成: 成功 {success_count} 只,失败 {fail_count} 只")
def get_all_stock_codes(self) -> List[str]:
"""
获取数据库中所有股票代码
Returns:
List[str]: 股票代码列表(不含后缀 .SZ/.SH)
"""
ts_codes = self.db.get_all_stock_codes()
return [code.split('.')[0] for code in ts_codes]
def delete_stock_data(self, code: str) -> bool:
"""
删除指定股票的所有数据
Args:
code: 股票代码(如 '000001' 或 '000001.SZ')
Returns:
bool: 是否删除成功
"""
ts_code = self.format_ts_code(code)
return self.db.delete_stock_data(ts_code)
def get_data_info(self) -> dict:
"""
获取数据信息
Returns:
dict: 数据信息
"""
info = self.db.get_database_info()
stock_codes = self.get_all_stock_codes()
info['stock_count'] = len(stock_codes)
info['stock_list'] = stock_codes
return info