mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 16:44:15 +08:00
fix(fuyao): 流式同步历史日K避免内存溢出
This commit is contained in:
@@ -32,12 +32,13 @@ import logging
|
|||||||
import math
|
import math
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable, Iterator
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import UTC, date, datetime, timedelta
|
from datetime import UTC, date, datetime, timedelta
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import polars as pl
|
import polars as pl
|
||||||
|
import pyarrow.parquet as pq
|
||||||
|
|
||||||
from app.data_providers.normalizer import DAILY_COLS, normalize_daily
|
from app.data_providers.normalizer import DAILY_COLS, normalize_daily
|
||||||
from app.indicators.pipeline import filter_halt_days
|
from app.indicators.pipeline import filter_halt_days
|
||||||
@@ -69,6 +70,12 @@ _DAILY10_DUMP_KIND = "daily-k-10d"
|
|||||||
_DAILY_DUMP_KIND = "daily-k" # 10 年全量日K dump(约 172MB), 深窗口一次下载覆盖全市场
|
_DAILY_DUMP_KIND = "daily-k" # 10 年全量日K dump(约 172MB), 深窗口一次下载覆盖全市场
|
||||||
_RECENT_DUMP_DAYS = 12 # 窗口跨度 ≤ 此天数时优先走 10d dump(覆盖 ≈10 个交易日)
|
_RECENT_DUMP_DAYS = 12 # 窗口跨度 ≤ 此天数时优先走 10d dump(覆盖 ≈10 个交易日)
|
||||||
_PREV_CLOSE_BACKDAYS = 30 # 推导因子时向前找"除权日前收盘"的回看天数(容忍长期停牌)
|
_PREV_CLOSE_BACKDAYS = 30 # 推导因子时向前找"除权日前收盘"的回看天数(容忍长期停牌)
|
||||||
|
_DAILY_DUMP_BATCH_ROWS = 100_000
|
||||||
|
_HIST_SYMBOL_BATCH = 50
|
||||||
|
_DAILY_DUMP_COLUMNS = [
|
||||||
|
"thscode", "adjusted", "date_ms", "open_price", "high_price", "low_price",
|
||||||
|
"close_price", "volume", "turnover",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def get_api_key() -> str:
|
def get_api_key() -> str:
|
||||||
@@ -444,34 +451,152 @@ class FuyaoProvider:
|
|||||||
- 兜底: 单标的 historical 接口(窗口早于 dump 覆盖 / dump 不可用; 10 年自动分片,
|
- 兜底: 单标的 historical 接口(窗口早于 dump 覆盖 / dump 不可用; 10 年自动分片,
|
||||||
逐标的节流 + 进度回调)。
|
逐标的节流 + 进度回调)。
|
||||||
"""
|
"""
|
||||||
|
chunks = list(self.iter_daily(
|
||||||
|
symbols,
|
||||||
|
start_time=start_time,
|
||||||
|
end_time=end_time,
|
||||||
|
asset_type=asset_type,
|
||||||
|
on_chunk_done=on_chunk_done,
|
||||||
|
))
|
||||||
|
return pl.concat(chunks, how="diagonal_relaxed") if chunks else pl.DataFrame()
|
||||||
|
|
||||||
|
def iter_daily(
|
||||||
|
self,
|
||||||
|
symbols: list[str],
|
||||||
|
start_time: datetime | None,
|
||||||
|
end_time: datetime | None,
|
||||||
|
asset_type: str = "stock",
|
||||||
|
on_chunk_done: Callable[[int, int], None] | None = None,
|
||||||
|
strict: bool = False,
|
||||||
|
) -> Iterator[pl.DataFrame]:
|
||||||
|
"""分批产出日K,供历史同步逐批落盘,避免全市场结果累积在内存。"""
|
||||||
if not symbols or asset_type != "stock":
|
if not symbols or asset_type != "stock":
|
||||||
return pl.DataFrame()
|
return
|
||||||
end_dt = end_time or datetime.now()
|
end_dt = end_time or datetime.now()
|
||||||
start_dt = start_time or (end_dt - timedelta(days=365))
|
start_dt = start_time or (end_dt - timedelta(days=365))
|
||||||
start_d, end_d = start_dt.date(), end_dt.date()
|
start_d, end_d = start_dt.date(), end_dt.date()
|
||||||
|
symset = set(symbols)
|
||||||
|
failed_symbols: list[str] = []
|
||||||
|
|
||||||
if (end_d - start_d).days <= _RECENT_DUMP_DAYS:
|
if (end_d - start_d).days <= _RECENT_DUMP_DAYS:
|
||||||
try:
|
try:
|
||||||
dump = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
dump = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||||
if _dump_covers(dump, start_d, end_d):
|
if _dump_covers(dump, start_d, end_d):
|
||||||
df = self._daily_from_dump(dump, set(symbols), start_d, end_d)
|
df = self._daily_from_dump(dump, symset, start_d, end_d)
|
||||||
if on_chunk_done:
|
if on_chunk_done:
|
||||||
on_chunk_done(1, 1)
|
on_chunk_done(1, 1)
|
||||||
logger.info("扶摇日K(10d dump)完成: %d 行 [%s ~ %s]", df.height, start_d, end_d)
|
if not df.is_empty():
|
||||||
return df
|
yield df
|
||||||
logger.info("扶摇 10d dump 未覆盖窗口 [%s ~ %s], 尝试 10 年 dump", start_d, end_d)
|
return
|
||||||
except FuyaoError as e:
|
except FuyaoError as e:
|
||||||
logger.warning("扶摇日K 10d dump 不可用, 尝试 10 年 dump: %s", e)
|
logger.warning("扶摇 10d dump 不可用: %s", e)
|
||||||
|
|
||||||
df = self._daily_from_big_dump(set(symbols), start_d, end_d)
|
dump_info = self._daily_dump_info()
|
||||||
if df is not None:
|
sources: list[tuple[str, date, date]] = []
|
||||||
if on_chunk_done:
|
if dump_info:
|
||||||
on_chunk_done(1, 1)
|
_, dump_min, dump_max = dump_info
|
||||||
logger.info("扶摇日K(10 年 dump)完成: %d 行 [%s ~ %s]", df.height, start_d, end_d)
|
if start_d < dump_min:
|
||||||
return df
|
sources.append(("api", start_d, min(end_d, dump_min - timedelta(days=1))))
|
||||||
df = self._daily_from_api(symbols, start_d, end_d, on_chunk_done)
|
overlap_start, overlap_end = max(start_d, dump_min), min(end_d, dump_max)
|
||||||
logger.info("扶摇日K(单标的接口)完成: %d 行 [%s ~ %s]", df.height, start_d, end_d)
|
if overlap_start <= overlap_end:
|
||||||
return df
|
sources.append(("dump", overlap_start, overlap_end))
|
||||||
|
tail_start = max(start_d, dump_max + timedelta(days=1))
|
||||||
|
if tail_start <= end_d:
|
||||||
|
try:
|
||||||
|
ten = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||||
|
ten_dates = pl.from_epoch(
|
||||||
|
ten["date_ms"].cast(pl.Int64) + _SH_MS, time_unit="ms"
|
||||||
|
).dt.date()
|
||||||
|
ten_min, ten_max = ten_dates.min(), ten_dates.max()
|
||||||
|
except (FuyaoError, KeyError):
|
||||||
|
ten_min = ten_max = None
|
||||||
|
if (
|
||||||
|
ten_min is not None
|
||||||
|
and ten_min <= dump_max + timedelta(days=1)
|
||||||
|
and _tail_ok(end_d, ten_max)
|
||||||
|
):
|
||||||
|
sources.append(("10d", tail_start, end_d))
|
||||||
|
else:
|
||||||
|
# 多年 dump 与请求终点之间存在不可验证的缺口,不能返回半段数据。
|
||||||
|
sources = [("api", start_d, end_d)]
|
||||||
|
else:
|
||||||
|
sources.append(("api", start_d, end_d))
|
||||||
|
|
||||||
|
dump_batch_count = 0
|
||||||
|
if dump_info:
|
||||||
|
dump_rows = pq.ParquetFile(dump_info[0]).metadata.num_rows
|
||||||
|
dump_batch_count = max(
|
||||||
|
1, (dump_rows + _DAILY_DUMP_BATCH_ROWS - 1) // _DAILY_DUMP_BATCH_ROWS
|
||||||
|
)
|
||||||
|
api_batch_count = (len(symbols) + _HIST_SYMBOL_BATCH - 1) // _HIST_SYMBOL_BATCH
|
||||||
|
total = sum(
|
||||||
|
dump_batch_count if kind == "dump" else 1 if kind == "10d" else api_batch_count
|
||||||
|
for kind, _, _ in sources
|
||||||
|
)
|
||||||
|
done = 0
|
||||||
|
for kind, source_start, source_end in sources:
|
||||||
|
if source_start > source_end:
|
||||||
|
continue
|
||||||
|
if kind == "dump":
|
||||||
|
path = dump_info[0] # type: ignore[index]
|
||||||
|
for df in self._iter_big_dump(path, symset, source_start, source_end):
|
||||||
|
done += 1
|
||||||
|
if on_chunk_done:
|
||||||
|
on_chunk_done(done, total)
|
||||||
|
if not df.is_empty():
|
||||||
|
yield df
|
||||||
|
elif kind == "10d":
|
||||||
|
ten = self._ensure_dump(_DAILY10_DUMP_KIND, "daily_k_10d")
|
||||||
|
df = self._daily_from_dump(ten, symset, source_start, source_end)
|
||||||
|
done += 1
|
||||||
|
if on_chunk_done:
|
||||||
|
on_chunk_done(done, total)
|
||||||
|
if not df.is_empty():
|
||||||
|
yield df
|
||||||
|
else:
|
||||||
|
batches = [
|
||||||
|
symbols[i:i + _HIST_SYMBOL_BATCH]
|
||||||
|
for i in range(0, len(symbols), _HIST_SYMBOL_BATCH)
|
||||||
|
]
|
||||||
|
for batch in batches:
|
||||||
|
rows: list[dict] = []
|
||||||
|
for symbol in batch:
|
||||||
|
rows.extend(_kline_rows(
|
||||||
|
symbol,
|
||||||
|
self._historical_bars(
|
||||||
|
symbol, source_start, source_end, failed_symbols=failed_symbols
|
||||||
|
),
|
||||||
|
))
|
||||||
|
time.sleep(_HIST_INTERVAL_S)
|
||||||
|
df = normalize_daily(rows, source=self.name)
|
||||||
|
done += 1
|
||||||
|
if on_chunk_done:
|
||||||
|
on_chunk_done(done, total)
|
||||||
|
if not df.is_empty():
|
||||||
|
yield df
|
||||||
|
if strict and failed_symbols:
|
||||||
|
sample = ", ".join(failed_symbols[:10])
|
||||||
|
raise FuyaoError(
|
||||||
|
f"扶摇日K同步部分标的失败: {len(failed_symbols)} 只 (样例: {sample})"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _daily_dump_info(self) -> tuple[Path, date, date] | None:
|
||||||
|
"""返回多年 dump 的路径和覆盖范围,不把大文件读入进程内存。"""
|
||||||
|
path = None
|
||||||
|
for candidate in sorted(_cache_dir().glob("daily_k__*.parquet"), reverse=True):
|
||||||
|
try:
|
||||||
|
dmin, dmax = _dump_date_range(candidate)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
if dmin is not None and dmax is not None:
|
||||||
|
return candidate, dmin, dmax
|
||||||
|
try:
|
||||||
|
path = self._ensure_dump_path(_DAILY_DUMP_KIND, "daily_k")
|
||||||
|
dmin, dmax = _dump_date_range(path)
|
||||||
|
except FuyaoError as e:
|
||||||
|
logger.warning("扶摇 10 年 dump 不可用, 回退单标的接口: %s", e)
|
||||||
|
return None
|
||||||
|
return (path, dmin, dmax) if dmin is not None and dmax is not None else None
|
||||||
|
|
||||||
def _daily_from_dump(
|
def _daily_from_dump(
|
||||||
self, dump: pl.DataFrame, symset: set[str], start_d: date, end_d: date
|
self, dump: pl.DataFrame, symset: set[str], start_d: date, end_d: date
|
||||||
@@ -481,6 +606,32 @@ class FuyaoProvider:
|
|||||||
)
|
)
|
||||||
return self._map_daily_dump(df, symset, start_d, end_d)
|
return self._map_daily_dump(df, symset, start_d, end_d)
|
||||||
|
|
||||||
|
def _iter_big_dump(
|
||||||
|
self, path: Path, symset: set[str], start_d: date, end_d: date
|
||||||
|
) -> Iterator[pl.DataFrame]:
|
||||||
|
"""按固定 record batch 读取多年 dump,不做单次全量 collect。"""
|
||||||
|
parquet = pq.ParquetFile(path)
|
||||||
|
columns = [name for name in _DAILY_DUMP_COLUMNS if name in parquet.schema.names]
|
||||||
|
for batch in parquet.iter_batches(
|
||||||
|
batch_size=_DAILY_DUMP_BATCH_ROWS,
|
||||||
|
columns=columns,
|
||||||
|
):
|
||||||
|
raw = pl.from_arrow(batch)
|
||||||
|
if raw.is_empty() or "date_ms" not in raw.columns or "thscode" not in raw.columns:
|
||||||
|
continue
|
||||||
|
start_ms, end_ms = _ms_of_date(start_d), _ms_of_date(end_d)
|
||||||
|
raw = raw.filter(
|
||||||
|
(pl.col("date_ms") >= start_ms)
|
||||||
|
& (pl.col("date_ms") <= end_ms)
|
||||||
|
& pl.col("thscode").is_in(sorted(symset))
|
||||||
|
)
|
||||||
|
if raw.is_empty():
|
||||||
|
continue
|
||||||
|
raw = raw.with_columns(
|
||||||
|
pl.from_epoch(pl.col("date_ms") + _SH_MS, time_unit="ms").dt.date().alias("date")
|
||||||
|
)
|
||||||
|
yield self._map_daily_dump(raw, symset, start_d, end_d)
|
||||||
|
|
||||||
def _map_daily_dump(
|
def _map_daily_dump(
|
||||||
self, df: pl.DataFrame, symset: set[str], start_d: date, end_d: date
|
self, df: pl.DataFrame, symset: set[str], start_d: date, end_d: date
|
||||||
) -> pl.DataFrame:
|
) -> pl.DataFrame:
|
||||||
@@ -582,7 +733,13 @@ class FuyaoProvider:
|
|||||||
on_chunk_done(i + 1, len(symbols))
|
on_chunk_done(i + 1, len(symbols))
|
||||||
return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame()
|
return pl.concat(frames, how="diagonal_relaxed") if frames else pl.DataFrame()
|
||||||
|
|
||||||
def _historical_bars(self, symbol: str, start_d: date, end_d: date) -> list[dict]:
|
def _historical_bars(
|
||||||
|
self,
|
||||||
|
symbol: str,
|
||||||
|
start_d: date,
|
||||||
|
end_d: date,
|
||||||
|
failed_symbols: list[str] | None = None,
|
||||||
|
) -> list[dict]:
|
||||||
"""按 ≤10 年窗口分片拉取单标的原始日K。中途失败软返回已得行, 不抛出。"""
|
"""按 ≤10 年窗口分片拉取单标的原始日K。中途失败软返回已得行, 不抛出。"""
|
||||||
out: list[dict] = []
|
out: list[dict] = []
|
||||||
s = _ms_of_date(start_d)
|
s = _ms_of_date(start_d)
|
||||||
@@ -593,6 +750,8 @@ class FuyaoProvider:
|
|||||||
out.extend(self._get_client().historical_kline(symbol, s, chunk_end, adjust="none"))
|
out.extend(self._get_client().historical_kline(symbol, s, chunk_end, adjust="none"))
|
||||||
except FuyaoError as err:
|
except FuyaoError as err:
|
||||||
logger.warning("扶摇日K拉取失败 %s [%s ~ %s]: %s", symbol, start_d, end_d, err)
|
logger.warning("扶摇日K拉取失败 %s [%s ~ %s]: %s", symbol, start_d, end_d, err)
|
||||||
|
if failed_symbols is not None:
|
||||||
|
failed_symbols.append(symbol)
|
||||||
break
|
break
|
||||||
if s + _HIST_MAX_SPAN_MS <= e:
|
if s + _HIST_MAX_SPAN_MS <= e:
|
||||||
time.sleep(_HIST_INTERVAL_S)
|
time.sleep(_HIST_INTERVAL_S)
|
||||||
|
|||||||
@@ -8,6 +8,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import shutil
|
||||||
|
import uuid
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from datetime import date, datetime, timedelta
|
from datetime import date, datetime, timedelta
|
||||||
|
|
||||||
@@ -177,6 +179,18 @@ def sync_and_persist_daily_batch(
|
|||||||
end_time = end_date or datetime.now()
|
end_time = end_date or datetime.now()
|
||||||
days = count or 365
|
days = count or 365
|
||||||
start_time = start_date or (end_time - timedelta(days=days))
|
start_time = start_date or (end_time - timedelta(days=days))
|
||||||
|
iter_daily = getattr(provider, "iter_daily", None)
|
||||||
|
if callable(iter_daily):
|
||||||
|
return _persist_daily_chunks(
|
||||||
|
iter_daily(
|
||||||
|
symbols,
|
||||||
|
start_time=start_time,
|
||||||
|
end_time=end_time,
|
||||||
|
on_chunk_done=on_chunk_done,
|
||||||
|
strict=True,
|
||||||
|
),
|
||||||
|
repo,
|
||||||
|
)
|
||||||
df = provider.get_daily(
|
df = provider.get_daily(
|
||||||
symbols,
|
symbols,
|
||||||
start_time=start_time,
|
start_time=start_time,
|
||||||
@@ -228,6 +242,45 @@ def sync_and_persist_daily_batch(
|
|||||||
return df.height
|
return df.height
|
||||||
|
|
||||||
|
|
||||||
|
def _persist_daily_chunks(chunks, repo: KlineRepository) -> int:
|
||||||
|
"""先把流式 provider 结果写入私有 staging,完整取数后再提交正式分区。"""
|
||||||
|
root = repo.store.data_dir / ".daily_sync_staging" / uuid.uuid4().hex
|
||||||
|
written = 0
|
||||||
|
try:
|
||||||
|
for index, df in enumerate(chunks):
|
||||||
|
if df.is_empty():
|
||||||
|
continue
|
||||||
|
for date_df in df.partition_by("date"):
|
||||||
|
dt = date_df["date"][0]
|
||||||
|
ds = dt.isoformat() if hasattr(dt, "isoformat") else str(dt)
|
||||||
|
out = root / f"date={ds}" / f"part-{index}.parquet"
|
||||||
|
out.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
date_df.write_parquet(out)
|
||||||
|
written += date_df.height
|
||||||
|
|
||||||
|
for date_dir in sorted(root.glob("date=*")):
|
||||||
|
files = sorted(date_dir.glob("*.parquet"))
|
||||||
|
if files:
|
||||||
|
repo.append_daily(pl.scan_parquet(files).collect(streaming=True))
|
||||||
|
finally:
|
||||||
|
shutil.rmtree(root, ignore_errors=True)
|
||||||
|
if root.parent.exists() and not any(root.parent.iterdir()):
|
||||||
|
root.parent.rmdir()
|
||||||
|
|
||||||
|
if written:
|
||||||
|
try:
|
||||||
|
d = repo.store.data_dir.as_posix()
|
||||||
|
repo.db.execute(
|
||||||
|
f"""CREATE OR REPLACE VIEW kline_daily AS
|
||||||
|
SELECT * FROM read_parquet(
|
||||||
|
'{d}/kline_daily/**/*.parquet', union_by_name=true
|
||||||
|
)"""
|
||||||
|
)
|
||||||
|
except Exception as e: # noqa: BLE001
|
||||||
|
logger.warning("refresh view failed: %s", e)
|
||||||
|
return written
|
||||||
|
|
||||||
|
|
||||||
def sync_daily_by_quotes(repo: KlineRepository) -> int:
|
def sync_daily_by_quotes(repo: KlineRepository) -> int:
|
||||||
"""用实时行情接口拉全市场当日数据,覆写 kline_daily 今天分区。
|
"""用实时行情接口拉全市场当日数据,覆写 kline_daily 今天分区。
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
"""扶摇流式日K同步: 大历史不得在内存中累积后再写入。"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import date, datetime
|
||||||
|
|
||||||
|
import polars as pl
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from app.services import kline_sync
|
||||||
|
from app.tickflow.repository import DataStore, KlineRepository
|
||||||
|
|
||||||
|
|
||||||
|
def _daily(symbol: str, day: date) -> pl.DataFrame:
|
||||||
|
return pl.DataFrame({
|
||||||
|
"symbol": [symbol], "date": [day], "open": [10.0], "high": [11.0],
|
||||||
|
"low": [9.0], "close": [10.5], "volume": [100.0], "amount": [1050.0],
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
class _StreamingProvider:
|
||||||
|
def __init__(self, chunks, error: Exception | None = None):
|
||||||
|
self.chunks = chunks
|
||||||
|
self.error = error
|
||||||
|
self.get_daily_called = False
|
||||||
|
|
||||||
|
def iter_daily(self, *args, **kwargs):
|
||||||
|
yield from self.chunks
|
||||||
|
if self.error:
|
||||||
|
raise self.error
|
||||||
|
|
||||||
|
def get_daily(self, *args, **kwargs):
|
||||||
|
self.get_daily_called = True
|
||||||
|
raise AssertionError("streaming provider must not collect a full daily DataFrame")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def repo(tmp_path):
|
||||||
|
return KlineRepository(DataStore(tmp_path))
|
||||||
|
|
||||||
|
|
||||||
|
def _route_fuyao(monkeypatch, provider):
|
||||||
|
monkeypatch.setattr(kline_sync.preferences, "get_daily_data_provider", lambda: "fuyao")
|
||||||
|
from app.data_providers import custom
|
||||||
|
monkeypatch.setattr(
|
||||||
|
custom,
|
||||||
|
"provider_has_dataset",
|
||||||
|
lambda name, dataset: name == "fuyao" and dataset == "daily",
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(custom, "get_provider", lambda name: provider)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_persists_streamed_fuyao_chunks_only_after_fetch(monkeypatch, repo):
|
||||||
|
provider = _StreamingProvider([
|
||||||
|
_daily("000001.SZ", date(2020, 1, 2)),
|
||||||
|
_daily("000002.SZ", date(2020, 1, 2)),
|
||||||
|
])
|
||||||
|
_route_fuyao(monkeypatch, provider)
|
||||||
|
|
||||||
|
written = kline_sync.sync_and_persist_daily_batch(
|
||||||
|
["000001.SZ", "000002.SZ"], repo, object(),
|
||||||
|
start_date=datetime(2020, 1, 1), end_date=datetime(2020, 1, 2),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert written == 2
|
||||||
|
assert provider.get_daily_called is False
|
||||||
|
stored = pl.read_parquet(
|
||||||
|
repo.store.data_dir / "kline_daily" / "date=2020-01-02" / "part.parquet"
|
||||||
|
)
|
||||||
|
assert set(stored["symbol"].to_list()) == {"000001.SZ", "000002.SZ"}
|
||||||
|
assert not (repo.store.data_dir / ".daily_sync_staging").exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_discards_staging_when_streaming_fails(monkeypatch, repo):
|
||||||
|
provider = _StreamingProvider(
|
||||||
|
[_daily("000001.SZ", date(2020, 1, 2))], error=RuntimeError("network lost")
|
||||||
|
)
|
||||||
|
_route_fuyao(monkeypatch, provider)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="network lost"):
|
||||||
|
kline_sync.sync_and_persist_daily_batch(
|
||||||
|
["000001.SZ"], repo, object(),
|
||||||
|
start_date=datetime(2020, 1, 1), end_date=datetime(2020, 1, 2),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert not list((repo.store.data_dir / "kline_daily").glob("date=*"))
|
||||||
|
assert not (repo.store.data_dir / ".daily_sync_staging").exists()
|
||||||
@@ -740,6 +740,39 @@ def test_daily_api_soft_fail_per_symbol(monkeypatch):
|
|||||||
assert df["symbol"].unique().to_list() == ["000001.SZ"]
|
assert df["symbol"].unique().to_list() == ["000001.SZ"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_iter_daily_strict_raises_after_api_failure(monkeypatch):
|
||||||
|
provider = _hist_provider(
|
||||||
|
monkeypatch,
|
||||||
|
_FakeHistClient(
|
||||||
|
{"000001.SZ": [_bar(date(2018, 1, 2), 10.0)]},
|
||||||
|
error_syms=("000002.SZ",),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
with pytest.raises(fc.FuyaoError, match="部分标的失败"):
|
||||||
|
list(provider.iter_daily(
|
||||||
|
["000001.SZ", "000002.SZ"],
|
||||||
|
datetime(2018, 1, 1), datetime(2018, 1, 3),
|
||||||
|
strict=True,
|
||||||
|
))
|
||||||
|
|
||||||
|
|
||||||
|
def test_iter_daily_api_is_bounded_by_symbol_batch(monkeypatch):
|
||||||
|
"""历史 API 每批只产出固定 symbol 集合,不积累全市场 frames。"""
|
||||||
|
bars = {
|
||||||
|
"000001.SZ": [_bar(date(2018, 1, 2), 10.0)],
|
||||||
|
"000002.SZ": [_bar(date(2018, 1, 2), 11.0)],
|
||||||
|
"000003.SZ": [_bar(date(2018, 1, 2), 12.0)],
|
||||||
|
}
|
||||||
|
provider = _hist_provider(monkeypatch, _FakeHistClient(bars))
|
||||||
|
monkeypatch.setattr(fp, "_HIST_SYMBOL_BATCH", 2)
|
||||||
|
chunks = list(provider.iter_daily(
|
||||||
|
list(bars), datetime(2018, 1, 1), datetime(2018, 1, 3)
|
||||||
|
))
|
||||||
|
assert [set(df["symbol"].to_list()) for df in chunks] == [
|
||||||
|
{"000001.SZ", "000002.SZ"}, {"000003.SZ"}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_daily_empty_symbols_or_non_stock_returns_empty(monkeypatch):
|
def test_daily_empty_symbols_or_non_stock_returns_empty(monkeypatch):
|
||||||
provider = _hist_provider(monkeypatch, _FakeHistClient({}))
|
provider = _hist_provider(monkeypatch, _FakeHistClient({}))
|
||||||
assert provider.get_daily([], None, None).is_empty()
|
assert provider.get_daily([], None, None).is_empty()
|
||||||
@@ -899,6 +932,37 @@ def test_daily_deep_window_uses_big_dump(monkeypatch, tmp_path):
|
|||||||
assert df["volume"].to_list() == [975_701.0, 12_345.0, 12_345.0]
|
assert df["volume"].to_list() == [975_701.0, 12_345.0, 12_345.0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_iter_daily_big_dump_reads_bounded_record_batches(monkeypatch, tmp_path):
|
||||||
|
"""多年 dump 分批读取,不经 LazyFrame collect() 物化整份文件。"""
|
||||||
|
rows = [_dump_bar("000001.SZ", date(2026, 1, 1) + timedelta(days=i), 10.0 + i)
|
||||||
|
for i in range(5)]
|
||||||
|
provider = _bigdump_provider(monkeypatch, tmp_path, rows)
|
||||||
|
monkeypatch.setattr(fp, "_DAILY_DUMP_BATCH_ROWS", 2)
|
||||||
|
chunks = list(provider.iter_daily(
|
||||||
|
["000001.SZ"], datetime(2026, 1, 1), datetime(2026, 1, 5)
|
||||||
|
))
|
||||||
|
assert [chunk.height for chunk in chunks] == [2, 2, 1]
|
||||||
|
assert [d for chunk in chunks for d in chunk["date"].to_list()] == [
|
||||||
|
date(2026, 1, 1), date(2026, 1, 2), date(2026, 1, 3),
|
||||||
|
date(2026, 1, 4), date(2026, 1, 5),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_iter_daily_splits_older_history_from_big_dump(monkeypatch, tmp_path):
|
||||||
|
"""请求跨 dump 起点时,只把 dump 外那段回退到历史 API。"""
|
||||||
|
provider = _bigdump_provider(
|
||||||
|
monkeypatch,
|
||||||
|
tmp_path,
|
||||||
|
[_dump_bar("000001.SZ", date(2020, 1, 2), 10.0)],
|
||||||
|
)
|
||||||
|
provider._client = _FakeHistClient({"000001.SZ": [_bar(date(2019, 12, 31), 9.0)]})
|
||||||
|
rows = pl.concat(list(provider.iter_daily(
|
||||||
|
["000001.SZ"], datetime(2019, 12, 30), datetime(2020, 1, 2)
|
||||||
|
)))
|
||||||
|
assert rows["date"].to_list() == [date(2019, 12, 31), date(2020, 1, 2)]
|
||||||
|
assert len(provider._get_client().calls) == 1
|
||||||
|
|
||||||
|
|
||||||
def test_daily_big_dump_tail_filled_by_10d(monkeypatch, tmp_path):
|
def test_daily_big_dump_tail_filled_by_10d(monkeypatch, tmp_path):
|
||||||
"""大 dump 末端缺口(dmax 旧)由 10d dump 补尾, 两段拼接无缝。"""
|
"""大 dump 末端缺口(dmax 旧)由 10d dump 补尾, 两段拼接无缝。"""
|
||||||
big_rows = [
|
big_rows = [
|
||||||
|
|||||||
Reference in New Issue
Block a user