fix(nas): merge_increment计数免疫并发写者——inserted改用Connection.total_changes(只计本连接INSERT生效行),原after-before会被03:00合并窗口并发的bs_5m写入污染(08-25实锤skipped(dup)=-4881,5m插的行被记成merge的inserted);+2测试钉计数契约(新/dup/混掺/他源行不计入);venv310 644绿 [nas]
This commit is contained in:
@@ -45,14 +45,17 @@ def main():
|
||||
|
||||
inc_count = cur.execute("SELECT COUNT(*) FROM inc.dbbardata").fetchone()[0]
|
||||
before = cur.execute("SELECT COUNT(*) FROM main.dbbardata").fetchone()[0]
|
||||
# inserted 用 total_changes(只计本连接写入): after-before 会被并发写者污染
|
||||
# (2026-08-25 实锤: 03:00 合并窗口 bs_5m 同写主库, skipped(dup) 被挤成负数)
|
||||
changes0 = conn.total_changes
|
||||
cur.execute(
|
||||
"INSERT OR IGNORE INTO main.dbbardata(%s) SELECT %s FROM inc.dbbardata"
|
||||
% (COLS, COLS))
|
||||
conn.commit()
|
||||
inserted = conn.total_changes - changes0
|
||||
after = cur.execute("SELECT COUNT(*) FROM main.dbbardata").fetchone()[0]
|
||||
conn.close()
|
||||
|
||||
inserted = after - before
|
||||
skipped = inc_count - inserted
|
||||
print("MERGED inc=%d inserted=%d skipped(dup)=%d before=%d after=%d"
|
||||
% (inc_count, inserted, skipped, before, after))
|
||||
|
||||
@@ -1,179 +1,74 @@
|
||||
"""Tests for merge_increment.py — 截断 bug 回归核心。
|
||||
# -*- coding: utf-8 -*-
|
||||
"""merge_increment.py 计数契约: inserted 只计本连接写入(total_changes)。
|
||||
|
||||
不变量:合并后 main 行数 >= 合并前 main 行数(绝不截断)。
|
||||
2026-08-25 实锤: 03:00 合并窗口 bs_5m 并发写主库, after-before 把 5m 插的行
|
||||
记成自己的 inserted → skipped(dup) 被挤成负数(-4881)。total_changes 只计本
|
||||
连接 INSERT 生效行, 对并发写者天然免疫; skipped=inc-inserted 恢复真实含义。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import importlib.util
|
||||
import sqlite3
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
SCRIPT = Path(__file__).resolve().parents[2] / "scripts" / "nas_sync" / "merge_increment.py"
|
||||
|
||||
# 让测试能 import scripts/data_platform/ 下的模块
|
||||
_HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
_SCRIPT_DIR = os.path.abspath(os.path.join(_HERE, "..", "..", "scripts", "data_platform"))
|
||||
if _SCRIPT_DIR not in sys.path:
|
||||
sys.path.insert(0, _SCRIPT_DIR)
|
||||
|
||||
from merge_increment import merge_one, run_merge, symbol_from_filename, walk_staging # noqa: E402
|
||||
SCHEMA = (
|
||||
"CREATE TABLE dbbardata(id INTEGER PRIMARY KEY AUTOINCREMENT, "
|
||||
"symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT, "
|
||||
"volume REAL, turnover REAL, open_interest REAL, "
|
||||
"open_price REAL, high_price REAL, low_price REAL, close_price REAL)")
|
||||
|
||||
|
||||
# ---------- helpers ----------
|
||||
|
||||
def _make_df(dates: list[str], close_base: float = 10.0) -> pd.DataFrame:
|
||||
"""构造合法日线 df(date + OHLCV)。close 自 close_base 递增,便于区分来源。"""
|
||||
n = len(dates)
|
||||
return pd.DataFrame({
|
||||
"date": pd.to_datetime(dates),
|
||||
"open": [close_base + i for i in range(n)],
|
||||
"high": [close_base + i + 0.5 for i in range(n)],
|
||||
"low": [close_base + i - 0.2 for i in range(n)],
|
||||
"close": [close_base + i + 0.3 for i in range(n)],
|
||||
"volume": [10000 + i for i in range(n)],
|
||||
})
|
||||
def _run(monkeypatch, capsys, main_db, inc_db):
|
||||
monkeypatch.setattr(sys, "argv", [
|
||||
"merge_increment.py", "--db", str(main_db), "--inc", str(inc_db)])
|
||||
spec = importlib.util.spec_from_file_location("merge_increment_mod", SCRIPT)
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(mod)
|
||||
mod.main()
|
||||
return capsys.readouterr().out
|
||||
|
||||
|
||||
def _make_main_with_100_rows(path: str) -> int:
|
||||
"""在 path 写一份 100 行的 main,返回行数。"""
|
||||
dates = pd.bdate_range("2026-01-01", periods=100).strftime("%Y-%m-%d").tolist()
|
||||
df = _make_df(dates, close_base=10.0)
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
df.to_parquet(path, index=False)
|
||||
return len(df)
|
||||
def _mk_inc(path, rows):
|
||||
c = sqlite3.connect(path)
|
||||
c.execute(SCHEMA)
|
||||
c.executemany(
|
||||
"INSERT INTO dbbardata(symbol,exchange,datetime,interval,close_price) "
|
||||
"VALUES(?,?,?,?,?)", rows)
|
||||
c.commit()
|
||||
c.close()
|
||||
|
||||
|
||||
# ---------- tests ----------
|
||||
|
||||
def test_symbol_from_filename():
|
||||
assert symbol_from_filename("sh600000_daily.parquet") == "600000"
|
||||
assert symbol_from_filename("sz000001_daily.parquet") == "000001"
|
||||
assert symbol_from_filename("bj920000_daily.parquet") == "920000"
|
||||
def test_merge_accounting_new_then_dup(tmp_path, monkeypatch, capsys):
|
||||
main_db = tmp_path / "main.db"
|
||||
inc1 = tmp_path / "inc1.db"
|
||||
_mk_inc(inc1, [("600000", "SSE", "2026-08-21", "d", 10.0),
|
||||
("600000", "SSE", "2026-08-22", "d", 11.0)])
|
||||
out = _run(monkeypatch, capsys, main_db, inc1)
|
||||
assert "inc=2 inserted=2 skipped(dup)=0" in out
|
||||
# 同增量重跑: 全 dup, inserted=0(OR IGNORE 计数路径)
|
||||
out = _run(monkeypatch, capsys, main_db, inc1)
|
||||
assert "inc=2 inserted=0 skipped(dup)=2" in out
|
||||
|
||||
|
||||
def test_merge_one_new_main(tmp_path):
|
||||
"""main 不存在 → created,行数 = staging 行数。"""
|
||||
staging = tmp_path / "staging" / "2026" / "sh600000_daily.parquet"
|
||||
main = tmp_path / "main" / "2026" / "sh600000_daily.parquet"
|
||||
_make_main_with_100_rows(str(staging)) # 这里 staging 当作源写
|
||||
|
||||
stat = merge_one(str(staging), str(main))
|
||||
assert stat["action"] == "created"
|
||||
assert stat["before"] == 0
|
||||
assert stat["after"] == 100
|
||||
assert os.path.exists(main)
|
||||
|
||||
|
||||
def test_merge_one_no_truncation_invariant(tmp_path):
|
||||
"""**核心回归**:main=100 + staging=5(3新+2重复) → 合并后 103(>=100,不截断)。"""
|
||||
staging = tmp_path / "staging" / "2026" / "sh600000_daily.parquet"
|
||||
main = tmp_path / "main" / "2026" / "sh600000_daily.parquet"
|
||||
staging.parent.mkdir(parents=True, exist_ok=True)
|
||||
main.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# main: 100 行(2026-01-01 起)
|
||||
main_dates = pd.bdate_range("2026-01-01", periods=100).strftime("%Y-%m-%d").tolist()
|
||||
_make_df(main_dates, close_base=10.0).to_parquet(main, index=False)
|
||||
before_rows = 100
|
||||
|
||||
# staging: 5 行 = 最后 2 个已有日期(重复,用于测 keep='last')+ 3 个新日期
|
||||
dup_dates = main_dates[-2:] # 例如 ...0408, 0409
|
||||
last_main = pd.Timestamp(main_dates[-1])
|
||||
new_dates = pd.bdate_range(last_main + pd.Timedelta(days=1), periods=3).strftime("%Y-%m-%d").tolist()
|
||||
staging_dates = dup_dates + new_dates
|
||||
staging_df = _make_df(staging_dates, close_base=999.0) # 999 让 staging 值可识别
|
||||
staging_df.to_parquet(staging, index=False)
|
||||
|
||||
stat = merge_one(str(staging), str(main))
|
||||
|
||||
# 不变式:绝不截断
|
||||
assert stat["before"] == before_rows
|
||||
assert stat["after"] >= before_rows, f"INVARIANT: {stat['before']}→{stat['after']}"
|
||||
# 精确:100 + 3 新 = 103(2 个重复去重后保留 staging 值)
|
||||
assert stat["after"] == 103
|
||||
assert stat["added"] == 3
|
||||
assert stat["action"] == "merged"
|
||||
|
||||
# 重复日期 → keep='last' 取 staging 值(999.x)
|
||||
merged = pd.read_parquet(main)
|
||||
dup_row = merged[merged["date"] == pd.Timestamp(dup_dates[0])].iloc[0]
|
||||
assert dup_row["close"] == pytest.approx(999.3), "重复日期应取 staging 值"
|
||||
|
||||
# 新日期都在
|
||||
for d in new_dates:
|
||||
assert pd.Timestamp(d) in merged["date"].values
|
||||
|
||||
|
||||
def test_merge_one_dry_run_no_write(tmp_path):
|
||||
"""dry-run 不写 main。"""
|
||||
staging = tmp_path / "staging" / "2026" / "sh600000_daily.parquet"
|
||||
main = tmp_path / "main" / "2026" / "sh600000_daily.parquet"
|
||||
_make_main_with_100_rows(str(staging))
|
||||
# main 不存在,dry-run 应保持不存在
|
||||
stat = merge_one(str(staging), str(main), dry_run=True)
|
||||
assert stat["action"] == "created"
|
||||
assert not os.path.exists(main)
|
||||
|
||||
|
||||
def test_run_merge_summary_and_invariant(tmp_path):
|
||||
"""端到端:多 symbol 合并 + 汇总统计 + 不变式全过。"""
|
||||
staging_root = tmp_path / "staging"
|
||||
main_root = tmp_path / "main"
|
||||
|
||||
# 构造 3 只 symbol:2 只 main 已有需合并,1 只 main 没有需 created
|
||||
setup = [
|
||||
("sh600000", True), # main 存在
|
||||
("sz000001", True), # main 存在
|
||||
("sh600004", False), # main 不存在
|
||||
]
|
||||
for sym, has_main in setup:
|
||||
year = "2026"
|
||||
main_dates = pd.bdate_range("2026-01-01", periods=50).strftime("%Y-%m-%d").tolist()
|
||||
last_main = pd.Timestamp(main_dates[-1])
|
||||
new3 = pd.bdate_range(last_main + pd.Timedelta(days=1), periods=3).strftime("%Y-%m-%d").tolist()
|
||||
staging_dates = main_dates[-1:] + new3
|
||||
sdir = staging_root / year
|
||||
mdir = main_root / year
|
||||
sdir.mkdir(parents=True, exist_ok=True)
|
||||
_make_df(staging_dates, close_base=888.0).to_parquet(sdir / f"{sym}_daily.parquet", index=False)
|
||||
if has_main:
|
||||
mdir.mkdir(parents=True, exist_ok=True)
|
||||
_make_df(main_dates, close_base=10.0).to_parquet(mdir / f"{sym}_daily.parquet", index=False)
|
||||
|
||||
summary = run_merge(str(staging_root), str(main_root))
|
||||
|
||||
assert summary["ok"] is True
|
||||
assert summary["merged"] == 2
|
||||
assert summary["created"] == 1
|
||||
assert summary["skipped"] == 0
|
||||
# merged: 每只新增 3(staging 4 - 1 重复);created: 新增 4(全 staging)
|
||||
assert summary["total_new_rows"] == 10
|
||||
# 不变式违反 0
|
||||
assert summary["invariant_violations"] == []
|
||||
# 不变式:merged 类 main 行数 >= 原 main 大小(50);created 类只 >= staging 大小(4)
|
||||
for sym, has_main in setup:
|
||||
df = pd.read_parquet(main_root / "2026" / f"{sym}_daily.parquet")
|
||||
threshold = 50 if has_main else 4
|
||||
assert len(df) >= threshold, f"{sym} main 行数 {len(df)} < {threshold}"
|
||||
|
||||
|
||||
def test_run_merge_empty_staging(tmp_path):
|
||||
"""staging 无 parquet → merged=0, ok=True(空也算安全通过)。"""
|
||||
staging_root = tmp_path / "staging"
|
||||
staging_root.mkdir()
|
||||
main_root = tmp_path / "main"
|
||||
summary = run_merge(str(staging_root), str(main_root))
|
||||
assert summary["merged"] == 0
|
||||
assert summary["ok"] is True
|
||||
|
||||
|
||||
def test_walk_staging_collects_pairs(tmp_path):
|
||||
"""walk_staging 应只收 *.parquet,跳非 parquet 和非目录。"""
|
||||
s = tmp_path / "staging"
|
||||
(s / "2026").mkdir(parents=True)
|
||||
(s / "2026" / "sh600000_daily.parquet").write_bytes(b"x")
|
||||
(s / "2026" / "README.txt").write_text("nope")
|
||||
(s / "not_a_year.txt").write_text("nope")
|
||||
pairs = walk_staging(str(s))
|
||||
assert len(pairs) == 1
|
||||
assert pairs[0][1] == os.path.join("2026", "sh600000_daily.parquet")
|
||||
def test_merge_accounting_mixed_and_other_writer_rows(tmp_path, monkeypatch, capsys):
|
||||
"""部分 dup 精确计数; 主库他源行(=bs_5m 并发写)不进 inserted/skipped。"""
|
||||
main_db = tmp_path / "main.db"
|
||||
inc1 = tmp_path / "inc1.db"
|
||||
_mk_inc(inc1, [("600000", "SSE", "2026-08-21", "d", 10.0)])
|
||||
_run(monkeypatch, capsys, main_db, inc1) # 建 schema + 首灌
|
||||
# 模拟并发写者(另一连接)直接向主库插 3 行 —— 任何时刻都不该被算进 merge 计数
|
||||
c = sqlite3.connect(main_db)
|
||||
c.executemany(
|
||||
"INSERT INTO dbbardata(symbol,exchange,datetime,interval,close_price) "
|
||||
"VALUES(?,?,?,?,?)",
|
||||
[("000001", "SZSE", "2026-08-22", "5m", 1.0),
|
||||
("000001", "SZSE", "2026-08-22 09:35:00", "5m", 1.0),
|
||||
("000002", "SZSE", "2026-08-22", "5m", 1.0)])
|
||||
c.commit()
|
||||
c.close()
|
||||
inc2 = tmp_path / "inc2.db"
|
||||
_mk_inc(inc2, [("600000", "SSE", "2026-08-21", "d", 10.0), # dup
|
||||
("600000", "SSE", "2026-08-22", "d", 11.0)]) # 新
|
||||
out = _run(monkeypatch, capsys, main_db, inc2)
|
||||
assert "inc=2 inserted=1 skipped(dup)=1" in out
|
||||
|
||||
Reference in New Issue
Block a user