data_cache.py

"""
数据缓存模块
用于管理数据的下载和本地缓存
"""

import logging
import os
import hashlib
from typing import Dict, Any
import pandas as pd
from datetime import datetime, timedelta


logger = logging.getLogger("AIQuant.data_cache")
logger.setLevel(logging.INFO)


class DataCache:
    """
    数据缓存管理器
    """
    def __init__(self, cache_dir: str = "data_cache"):
        """
        初始化数据缓存管理器

        Args:
            cache_dir: 缓存目录
        """
        self.cache_dir = cache_dir
        self._ensure_dir_exists()

    def _ensure_dir_exists(self):
        """
        确保缓存目录存在
        """
        if not os.path.exists(self.cache_dir):
            os.makedirs(self.cache_dir, exist_ok=True)
            logger.info(f"创建数据缓存目录: {self.cache_dir}")

    def _generate_cache_key(self,
                             symbol: str,
                             frequency: str = "d",
                             start_date: str = None,
                             end_date: str = None) -> str:
        """
        生成缓存键

        Args:
            symbol: 股票代码
            frequency: 频率
            start_date: 起始日期
            end_date: 结束日期

        Returns:
            缓存键
        """
        key = f"{symbol}_{frequency}"
        if start_date:
            key += f"_{start_date}"
        if end_date:
            key += f"_{end_date}"

        # 生成 MD5 哈希值作为文件名
        return hashlib.md5(key.encode()).hexdigest()

    def _get_cache_filepath(self, cache_key: str) -> str:
        """
        获取缓存文件路径

        Args:
            cache_key: 缓存键

        Returns:
            缓存文件路径
        """
        return os.path.join(self.cache_dir, f"{cache_key}.csv")

    def save(self,
             symbol: str,
             data: pd.DataFrame,
             frequency: str = "d",
             start_date: str = None,
             end_date: str = None) -> str:
        """
        保存数据到缓存

        Args:
            symbol: 股票代码
            data: 数据
            frequency: 频率
            start_date: 起始日期
            end_date: 结束日期

        Returns:
            缓存文件路径
        """
        self._ensure_dir_exists()

        cache_key = self._generate_cache_key(
            symbol, frequency, start_date, end_date)
        filepath = self._get_cache_filepath(cache_key)

        data.to_csv(filepath)
        logger.info(f"数据已保存到缓存: {filepath}")

        return filepath

    def load(self,
             symbol: str,
             frequency: str = "d",
             start_date: str = None,
             end_date: str = None) -> pd.DataFrame:
        """
        从缓存加载数据

        Args:
            symbol: 股票代码
            frequency: 频率
            start_date: 起始日期
            end_date: 结束日期

        Returns:
            加载的数据
        """
        cache_key = self._generate_cache_key(
            symbol, frequency, start_date, end_date)
        filepath = self._get_cache_filepath(cache_key)

        if os.path.exists(filepath):
            logger.info(f"从缓存加载数据: {filepath}")
            return pd.read_csv(filepath, index_col=0, parse_dates=True)
        else:
            logger.warning(f"缓存文件不存在: {filepath}")
            return pd.DataFrame()

    def exists(self,
               symbol: str,
               frequency: str = "d",
               start_date: str = None,
               end_date: str = None) -> bool:
        """
        检查缓存是否存在

        Args:
            symbol: 股票代码
            frequency: 频率
            start_date: 起始日期
            end_date: 结束日期

        Returns:
            缓存是否存在
        """
        cache_key = self._generate_cache_key(
            symbol, frequency, start_date, end_date)
        filepath = self._get_cache_filepath(cache_key)
        return os.path.exists(filepath)

    def delete(self,
               symbol: str,
               frequency: str = "d",
               start_date: str = None,
               end_date: str = None) -> bool:
        """
        删除缓存

        Args:
            symbol: 股票代码
            frequency: 频率
            start_date: 起始日期
            end_date: 结束日期

        Returns:
            删除是否成功
        """
        cache_key = self._generate_cache_key(
            symbol, frequency, start_date, end_date)
        filepath = self._get_cache_filepath(cache_key)

        if os.path.exists(filepath):
            os.remove(filepath)
            logger.info(f"缓存已删除: {filepath}")
            return True
        else:
            logger.warning(f"缓存文件不存在: {filepath}")
            return False

    def get_cache_info(self) -> Dict[str, Any]:
        """
        获取缓存信息

        Returns:
            缓存信息
        """
        if not os.path.exists(self.cache_dir):
            return {"cache_dir": self.cache_dir, "files": [], "total_size": 0}

        files = []
        total_size = 0

        for filename in os.listdir(self.cache_dir):
            if filename.endswith(".csv"):
                filepath = os.path.join(self.cache_dir, filename)
                file_size = os.path.getsize(filepath)
                file_mtime = datetime.fromtimestamp(os.path.getmtime(filepath))

                files.append({
                    "filename": filename,
                    "size": file_size,
                    "mtime": file_mtime
                })

                total_size += file_size

        files.sort(key=lambda x: x["mtime"], reverse=True)

        return {
            "cache_dir": self.cache_dir,
            "files": files,
            "total_size": total_size,
            "num_files": len(files),
        }

    def clear_expired(self, days: int = 30) -> int:
        """
        清理过期的缓存文件

        Args:
            days: 过期天数

        Returns:
            删除的文件数量
        """
        if not os.path.exists(self.cache_dir):
            return 0

        expired_date = datetime.now() - timedelta(days=days)
        deleted_count = 0

        for filename in os.listdir(self.cache_dir):
            if filename.endswith(".csv"):
                filepath = os.path.join(self.cache_dir, filename)
                file_mtime = datetime.fromtimestamp(os.path.getmtime(filepath))

                if file_mtime < expired_date:
                    os.remove(filepath)
                    deleted_count += 1
                    logger.info(f"已删除过期缓存: {filepath}")

        logger.info(f"共删除 {deleted_count} 个过期缓存文件")
        return deleted_count

    def clear_all(self) -> int:
        """
        清理所有缓存文件

        Returns:
            删除的文件数量
        """
        if not os.path.exists(self.cache_dir):
            return 0

        deleted_count = 0

        for filename in os.listdir(self.cache_dir):
            if filename.endswith(".csv"):
                filepath = os.path.join(self.cache_dir, filename)
                os.remove(filepath)
                deleted_count += 1
                logger.info(f"已删除缓存: {filepath}")

        logger.info(f"共删除 {deleted_count} 个缓存文件")
        return deleted_count


# 默认数据缓存管理器
default_cache = DataCache()


# 便捷函数
def save_data(symbol: str,
              data: pd.DataFrame,
              frequency: str = "d",
              start_date: str = None,
              end_date: str = None) -> str:
    """
    便捷函数:保存数据到缓存

    Args:
        symbol: 股票代码
        data: 数据
        frequency: 频率
        start_date: 起始日期
        end_date: 结束日期

    Returns:
        缓存文件路径
    """
    return default_cache.save(symbol, data, frequency, start_date, end_date)


def load_data(symbol: str,
              frequency: str = "d",
              start_date: str = None,
              end_date: str = None) -> pd.DataFrame:
    """
    便捷函数:从缓存加载数据

    Args:
        symbol: 股票代码
        frequency: 频率
        start_date: 起始日期
        end_date: 结束日期

    Returns:
        加载的数据
    """
    return default_cache.load(symbol, frequency, start_date, end_date)


def data_exists(symbol: str,
                frequency: str = "d",
                start_date: str = None,
                end_date: str = None) -> bool:
    """
    便捷函数:检查缓存是否存在

    Args:
        symbol: 股票代码
        frequency: 频率
        start_date: 起始日期
        end_date: 结束日期

    Returns:
        缓存是否存在
    """
    return default_cache.exists(symbol, frequency, start_date, end_date)


def delete_data(symbol: str,
                frequency: str = "d",
                start_date: str = None,
                end_date: str = None) -> bool:
    """
    便捷函数:删除缓存

    Args:
        symbol: 股票代码
        frequency: 频率
        start_date: 起始日期
        end_date: 结束日期

    Returns:
        删除是否成功
    """
    return default_cache.delete(symbol, frequency, start_date, end_date)


def get_cache_info() -> Dict[str, Any]:
    """
    便捷函数:获取缓存信息

    Returns:
        缓存信息
    """
    return default_cache.get_cache_info()


def clear_expired_cache(days: int = 30) -> int:
    """
    便捷函数:清理过期的缓存文件

    Args:
        days: 过期天数

    Returns:
        删除的文件数量
    """
    return default_cache.clear_expired(days)


def clear_all_cache() -> int:
    """
    便捷函数:清理所有缓存文件

    Returns:
        删除的文件数量
    """
    return default_cache.clear_all()


# 测试函数
if __name__ == "__main__":
    import sys
    import os

    # 添加项目根目录到路径
    sys.path.insert(0, os.path.dirname(
        os.path.dirname(os.path.abspath(__file__))))

    from service.data.mock_data import generate_random_kline_data

    # 测试数据缓存功能
    symbol = "000001.SZ"
    data = generate_random_kline_data(symbol)

    # 保存数据到缓存
    save_data(symbol, data)

    # 检查缓存是否存在
    print(f"缓存是否存在: {data_exists(symbol)}")

    # 从缓存加载数据
    loaded_data = load_data(symbol)
    print(f"加载的数据形状: {loaded_data.shape}")

    # 打印数据
    print(loaded_data.head())

    # 获取缓存信息
    cache_info = get_cache_info()
    print(f"缓存文件数量: {cache_info['num_files']}")
    print(f"缓存总大小: {cache_info['total_size']} 字节")

    # 删除数据
    delete_data(symbol)
    print(f"删除后缓存是否存在: {data_exists(symbol)}")