diff --git a/sanguo_data/datawriter.py b/sanguo_data/datawriter.py new file mode 100644 index 0000000..237abb0 --- /dev/null +++ b/sanguo_data/datawriter.py @@ -0,0 +1,34 @@ +# sanguo_data/datawriter.py +import os +import pandas as pd +from pathlib import Path +from vnpy.trader.object import BarData +from vnpy.trader.constant import Interval +from sanguo_data.datareader import _row_to_bar +from sanguo_data.config import DataConfig + +def atomic_write_parquet(path: str, df: pd.DataFrame) -> None: + p = Path(path) + p.parent.mkdir(parents=True, exist_ok=True) + tmp = str(p) + ".tmp" + df.to_parquet(tmp) + os.replace(tmp, str(p)) # 原子替换 + +def write_daily(symbol: str, df: pd.DataFrame, cfg: DataConfig) -> None: + # 1) parquet 增量合并(按年分区,去重保留最新) + for year, group in df.groupby(df["date"].str[:4]): + f = Path(cfg.data_paths["daily_dir"]) / year / f"{symbol}.parquet" + if f.exists(): + old = pd.read_parquet(f) + combined = pd.concat([old, group]).drop_duplicates("date", keep="last") + else: + combined = group + atomic_write_parquet(str(f), combined) + # 2) vnpy SQLite + bars = [_row_to_bar(symbol, row, Interval.DAILY) for _, row in df.iterrows()] + _save_to_vnpy_db(bars, cfg) + +def _save_to_vnpy_db(bars: list[BarData], cfg: DataConfig) -> None: + from vnpy.trader.database import get_database + db = get_database() + db.save_bar_data(bars) # spike 验证签名 diff --git a/tests/data/test_datawriter.py b/tests/data/test_datawriter.py new file mode 100644 index 0000000..58399e2 --- /dev/null +++ b/tests/data/test_datawriter.py @@ -0,0 +1,24 @@ +# tests/data/test_datawriter.py +import pandas as pd +from sanguo_data.config import DataConfig +from sanguo_data.datawriter import write_daily, atomic_write_parquet + +def test_atomic_write_parquet(tmp_path): + f = tmp_path / "2026" / "600000.parquet" + df = pd.DataFrame({"date": ["2026-01-01"], "open": [10.0]}) + atomic_write_parquet(str(f), df) + assert f.exists() + assert not list(tmp_path.glob("*.tmp")) + +def test_write_daily_writes_parquet_and_db(tmp_path, monkeypatch): + cfg = DataConfig( + data_paths={"daily_dir": str(tmp_path / "daily"), "vnpy_db": str(tmp_path / "q.db")}, + data_sources={}, validation={}, performance={}, + ) + df = pd.DataFrame({"date": ["2026-01-01"], "open": [10.0], "high": [10.0], + "low": [10.0], "close": [10.0], "volume": [100]}) + called = {} + monkeypatch.setattr("sanguo_data.datawriter._save_to_vnpy_db", lambda bars, cfg: called.setdefault("bars", bars)) + write_daily("600000", df, cfg) + assert (tmp_path / "daily" / "2026" / "600000.parquet").exists() + assert len(called["bars"]) == 1