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