feat(portfolio): LocalUnifiedProvider 批量行情接口 get_closes_panel (G5-Phase1)
单条 dbbardata 参数化查询 (symbol,exchange) OR + pivot 返宽表, 替代 N 次 get_price, 为策略向量化提速铺路。raw close 口径一致, 缺失 NaN 列, symbol-exchange 配对防歧义。22 tests 含 batch-vs-逐只回归。本文件另含 get_value_metrics 透传(策略移植)。
This commit is contained in:
@@ -241,6 +241,96 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
|
||||
return next(iter(frames.values()))
|
||||
return pd.concat(frames, axis=1)
|
||||
|
||||
# ==================== get_closes_panel (批量行情,G5 Phase1) ====================
|
||||
def get_closes_panel(
|
||||
self,
|
||||
symbols: List[str],
|
||||
start: Union[str, datetime],
|
||||
end: Union[str, datetime],
|
||||
interval: str = "d",
|
||||
) -> pd.DataFrame:
|
||||
"""批量取多只股票 close,返宽表 (index=datetime, columns=symbol, values=close)。
|
||||
|
||||
单条 dbbardata 参数化查询 ``(symbol, exchange) OR (...)`` + ``pivot``,
|
||||
替代选股时 N 次 ``get_price``(为 Phase 2 策略向量化提速 10-50× 铺路)。
|
||||
|
||||
与逐只 ``get_price(fq='raw')`` 等价(raw close, 不复权),区别仅在 IO 次数:
|
||||
- 逐只: N 次 SQL(N+1 query pattern)
|
||||
- 批量: 1 次 SQL + pivot
|
||||
|
||||
Args:
|
||||
symbols: jq 风格代码列表 (``["600519.XSHG", "000001.XSHE"]``) 或纯 6 位
|
||||
start/end: 日期 ``YYYY-MM-DD`` (或 datetime,取 date 部分)
|
||||
interval: ``"d"``=日线 (dbbardata ``interval`` 字段; ``"15m"`` 等 v2 扩展)
|
||||
|
||||
Returns:
|
||||
DataFrame, ``index=datetime`` (升序), ``columns=symbols``(按输入顺序),
|
||||
``values=close_price``。缺失股票 → 该列全 NaN; 重复行 dedup (keep last)。
|
||||
|
||||
空列表 → 空 DataFrame(index=DatetimeIndex,columns=[])。
|
||||
"""
|
||||
if not symbols:
|
||||
return pd.DataFrame(index=pd.DatetimeIndex([]))
|
||||
|
||||
# jq 代码 → (db_symbol, db_exchange); 保留 input_code 作列名
|
||||
# (symbol 不唯一: 932000 中证2000 / 北交所 920xxx 等, 必须按 exchange 配对)
|
||||
pairs: List[tuple[str, str, str]] = [] # (input_code, db_symbol, db_exchange)
|
||||
for s in symbols:
|
||||
sym, exc = jq_to_dbbardata(str(s))
|
||||
pairs.append((str(s), sym, exc))
|
||||
|
||||
conn = self._connect()
|
||||
start_str = self._to_date_str(start) or "1990-01-01"
|
||||
end_str = self._to_date_str(end) or datetime.now().strftime("%Y-%m-%d")
|
||||
|
||||
# 参数化 (symbol, exchange) OR 子句 — 防注入 + 不字符串拼接(长度无上限)
|
||||
where_parts = " OR ".join(["(symbol=? AND exchange=?)" for _ in pairs])
|
||||
params: List[Any] = [interval, start_str, end_str]
|
||||
for _inp, sym, exc in pairs:
|
||||
params.extend([sym, exc])
|
||||
|
||||
q = (
|
||||
"SELECT datetime, symbol, close_price FROM dbbardata "
|
||||
"WHERE interval=? "
|
||||
"AND substr(datetime,1,10)>=? AND substr(datetime,1,10)<=? "
|
||||
f"AND ({where_parts})"
|
||||
)
|
||||
df = pd.read_sql(q, conn, params=params)
|
||||
|
||||
if df.empty:
|
||||
# 全缺失: 返宽表骨架(全 NaN 列, 空 DatetimeIndex) — 与逐只 get_price 行为一致
|
||||
return pd.DataFrame(
|
||||
{inp: pd.Series(dtype=float) for inp in symbols},
|
||||
index=pd.DatetimeIndex([]),
|
||||
)
|
||||
|
||||
# 混合格式 datetime(与 get_price 一致: dbbardata 列混 "2024-09-26" 与
|
||||
# "2024-09-26 00:00:00",pandas 2.3 严格模式要 format="mixed")
|
||||
df["datetime"] = pd.to_datetime(df["datetime"], format="mixed")
|
||||
|
||||
# db_symbol → input_code 还原(列名回输入 jq code, 与策略其他接口一致)
|
||||
sym_to_input: Dict[str, str] = {}
|
||||
for inp, sym, _exc in pairs:
|
||||
sym_to_input.setdefault(sym, inp)
|
||||
df["symbol"] = df["symbol"].map(sym_to_input).fillna(df["symbol"])
|
||||
|
||||
# dedup: 同 (datetime, symbol) 重复行取最后一条(增量合并容错)
|
||||
df = df.drop_duplicates(subset=["datetime", "symbol"], keep="last")
|
||||
|
||||
# pivot 宽表
|
||||
wide = df.pivot(index="datetime", columns="symbol", values="close_price")
|
||||
|
||||
# 补缺失 symbol 列(全 NaN), 按 input 顺序对齐 columns
|
||||
for inp in symbols:
|
||||
if inp not in wide.columns:
|
||||
wide[inp] = float("nan")
|
||||
wide = wide[symbols]
|
||||
|
||||
# 升序 + 清 index name(与 get_price 一致)
|
||||
wide = wide.sort_index()
|
||||
wide.index.name = None
|
||||
return wide
|
||||
|
||||
# ==================== get_index_stocks (constituent_unified 并集) ====================
|
||||
def get_index_stocks(
|
||||
self,
|
||||
@@ -334,6 +424,18 @@ class LocalUnifiedProvider(DataProvider): # type: ignore[misc]
|
||||
df = df.set_index("code", drop=False)
|
||||
return df
|
||||
|
||||
def get_value_metrics(
|
||||
self,
|
||||
stock: str,
|
||||
date: Union[str, datetime],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""委托 LocalParquetProvider 读三表多期 + NOTICE_DATE 过滤(供 ValueSelectionStrategy)。
|
||||
|
||||
LocalParquetProvider 实现见 ``local_parquet_provider.get_value_metrics``;
|
||||
本类只是把职责转发给已有的 ``_lpp_helper``(DRY, 不复制字段映射逻辑)。
|
||||
"""
|
||||
return self._get_lpp_helper().get_value_metrics(stock, date)
|
||||
|
||||
def _build_fundamental_row(
|
||||
self, jq_code: str, date_str: str,
|
||||
) -> Dict[str, Any]:
|
||||
|
||||
Reference in New Issue
Block a user