- 新增 finance/tests/ 6 个测试套件(agents/backtest/dao_upsert/factors/features/fundamental_lookahead) - 数据层: data_manager / dao 优化,新增 upsert 逻辑 - 因子层: 基本面因子抽象定位 _mapping、ROE/PE/PB 重构 - 回测层: vectorbt/engine 大改动(251 行),report 增强 - ML 层: features/backtest_integration 特征工程与回测优化 - CLI: agent_cli 重构 - config/settings 扩充配置项
202 lines
8.0 KiB
Python
202 lines
8.0 KiB
Python
"""
|
|
DataManager — 统一数据管理层。
|
|
|
|
策略/模型层通过 DataManager 获取数据,不直接访问 AkShare/Tushare 或数据库。
|
|
优先从 DB 读取,缺失时依次尝试 AkShare → Tushare 拉取并入库。
|
|
"""
|
|
|
|
import logging
|
|
import time
|
|
|
|
import pandas as pd
|
|
|
|
from config.settings import DEFAULT_START_DATE, DEFAULT_END_DATE
|
|
from data.sources.akshare_source import AkShareSource
|
|
from data.sources.tushare_source import TushareSource
|
|
from database import dao
|
|
from database.models import create_all_tables
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _rollback_days(yyyymmdd: str, days: int) -> str:
|
|
"""把 'YYYYMMDD' 往前回退 days 个自然日,返回 'YYYYMMDD'。"""
|
|
from datetime import datetime, timedelta
|
|
dt = datetime.strptime(yyyymmdd, "%Y%m%d") - timedelta(days=days)
|
|
return dt.strftime("%Y%m%d")
|
|
|
|
|
|
def _safe_save_daily(df: pd.DataFrame, skip_exc: bool = True) -> None:
|
|
"""把拉取到的日线入库;失败时记日志(可选抛错)。"""
|
|
cols = [c for c in dao._DAILY_COLS if c in df.columns]
|
|
try:
|
|
dao.save_daily(df[cols])
|
|
except Exception as e:
|
|
if skip_exc:
|
|
logger.exception("[DataManager] 日线入库失败(数据已获取,未持久化):%s", e)
|
|
else:
|
|
raise
|
|
|
|
|
|
class DataManager:
|
|
"""统一数据管理。双数据源:AkShare(主)+ Tushare(备)。"""
|
|
|
|
def __init__(self):
|
|
self._ak_source: AkShareSource | None = None
|
|
self._ts_source: TushareSource | None = None
|
|
|
|
@property
|
|
def ak(self) -> AkShareSource:
|
|
if self._ak_source is None:
|
|
self._ak_source = AkShareSource()
|
|
return self._ak_source
|
|
|
|
@property
|
|
def ts(self) -> TushareSource:
|
|
if self._ts_source is None:
|
|
self._ts_source = TushareSource()
|
|
return self._ts_source
|
|
|
|
# ── 初始化 ────────────────────────────────────────────
|
|
|
|
def init_db(self) -> None:
|
|
create_all_tables()
|
|
|
|
# ── 内部:双源 try ────────────────────────────────────
|
|
|
|
def _try_fetch(self, method_name: str, *args, **kwargs):
|
|
"""
|
|
依次尝试 AkShare → Tushare 调用同一方法名。
|
|
|
|
method_name: 'fetch_daily' | 'fetch_stock_list' | 'fetch_financial'
|
|
返回: (result_df, source_name) 或 (empty_df, None)
|
|
|
|
如果 ts_code 是指数代码,自动路由到 fetch_index_daily。
|
|
"""
|
|
from data.sources.akshare_source import is_index_code
|
|
|
|
# 指数自动路由
|
|
if method_name == "fetch_daily" and args:
|
|
ts_code = args[0]
|
|
if is_index_code(ts_code):
|
|
method_name = "fetch_index_daily"
|
|
|
|
for label, source in [("Tushare", self.ts), ("AkShare", self.ak)]:
|
|
try:
|
|
if label == "Tushare" and not source.available:
|
|
continue
|
|
fn = getattr(source, method_name)
|
|
df = fn(*args, **kwargs)
|
|
if df is not None and not df.empty:
|
|
return df, label
|
|
except Exception as e:
|
|
logger.warning(" [%s] %s 失败: %s", label, method_name, e)
|
|
return pd.DataFrame(), None
|
|
|
|
# ── 股票列表 ──────────────────────────────────────────
|
|
|
|
def get_stock_list(self, force_refresh: bool = False) -> pd.DataFrame:
|
|
if not force_refresh:
|
|
df = dao.query_stock_list()
|
|
if not df.empty:
|
|
logger.info("[DataManager] 从 DB 读取股票列表: %s 只", len(df))
|
|
return df
|
|
|
|
logger.info("[DataManager] 拉取股票列表 (AkShare → Tushare)...")
|
|
df, src = self._try_fetch("fetch_stock_list")
|
|
if df.empty:
|
|
logger.warning("[DataManager] 所有数据源均无法获取股票列表")
|
|
return pd.DataFrame()
|
|
logger.info("[DataManager] 股票列表已入库 (%s): %s 只", src, len(df))
|
|
dao.save_stock_list(df)
|
|
time.sleep(2)
|
|
return df
|
|
|
|
# ── 日线数据 ──────────────────────────────────────────
|
|
|
|
def get_daily(
|
|
self,
|
|
ts_code: str,
|
|
start: str | None = None,
|
|
end: str | None = None,
|
|
force_refresh: bool = False,
|
|
) -> pd.DataFrame:
|
|
start = start or DEFAULT_START_DATE
|
|
end = end or DEFAULT_END_DATE or time.strftime("%Y%m%d")
|
|
|
|
if not force_refresh:
|
|
df = dao.query_daily(ts_code, start, end)
|
|
if not df.empty:
|
|
# 检查 DB 窗口尾部是否明显落后于请求的 end(>5 个自然日)。
|
|
# 若是,说明存在数据缺口,触发一次增量补拉,而不是把陈旧数据当完整返回。
|
|
latest = str(df["trade_date"].astype(str).max())
|
|
if end >= "20200101" and latest < _rollback_days(end, 7):
|
|
logger.warning(
|
|
"%s DB 窗口尾部 %s 明显落后于请求日 %s,触发补拉", ts_code, latest, end)
|
|
refreshed, src = self._try_fetch("fetch_daily", ts_code, latest, end)
|
|
if not refreshed.empty and src:
|
|
_safe_save_daily(refreshed)
|
|
return df
|
|
|
|
# DB 未命中 → 从数据源拉取
|
|
df, src = self._try_fetch("fetch_daily", ts_code, start, end)
|
|
if df.empty:
|
|
logger.warning("%s 日线获取失败 (AkShare+Tushare 均不可用)", ts_code)
|
|
return pd.DataFrame()
|
|
|
|
# 保存到 DB(失败仅记日志,不阻断已获取到的数据返回)
|
|
_safe_save_daily(df)
|
|
return df
|
|
|
|
def sync_daily(self, ts_code: str) -> int:
|
|
latest = dao.get_latest_trade_date(ts_code)
|
|
today = time.strftime("%Y%m%d")
|
|
if latest and latest >= today:
|
|
logger.info("[DataManager] %s 数据已是最新 (%s)", ts_code, latest)
|
|
return 0
|
|
|
|
start = latest or DEFAULT_START_DATE
|
|
df, src = self._try_fetch("fetch_daily", ts_code, start, today)
|
|
if df.empty:
|
|
logger.warning("[DataManager] %s sync_daily 失败 (AkShare+Tushare 均不可用)", ts_code)
|
|
return 0
|
|
|
|
# 入库失败要能反映到返回值,否则调用方会误以为已持久化
|
|
cols = [c for c in dao._DAILY_COLS if c in df.columns]
|
|
try:
|
|
dao.save_daily(df[cols])
|
|
except Exception as e:
|
|
logger.exception("[DataManager] %s 日线入库失败:%s", ts_code, e)
|
|
raise
|
|
logger.info("[DataManager] %s 同步 %s 条日线 (%s)", ts_code, len(df), src)
|
|
return len(df)
|
|
|
|
def sync_all_daily(self) -> int:
|
|
stock_list = self.get_stock_list()
|
|
total = 0
|
|
for i, ts_code in enumerate(stock_list.index):
|
|
try:
|
|
total += self.sync_daily(ts_code)
|
|
if (i + 1) % 50 == 0:
|
|
logger.info("[DataManager] 进度: %s/%s", i + 1, len(stock_list))
|
|
time.sleep(1)
|
|
except Exception as e:
|
|
logger.warning("[DataManager] %s 同步失败: %s", ts_code, e)
|
|
logger.info("[DataManager] 全量同步完成,新增 %s 条", total)
|
|
return total
|
|
|
|
# ── 财务数据 ──────────────────────────────────────────
|
|
|
|
def get_financial(self, ts_code: str) -> pd.DataFrame:
|
|
df = dao.query_financial(ts_code)
|
|
if not df.empty:
|
|
return df
|
|
df, src = self._try_fetch("fetch_financial", ts_code)
|
|
if not df.empty:
|
|
try:
|
|
cols = [c for c in dao._FINA_COLS if c in df.columns]
|
|
dao.save_financial(df[cols])
|
|
except Exception as e:
|
|
logger.exception("财务数据入库失败: %s", e)
|
|
return df
|