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)}")