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:
2026-08-25 21:32:28 +08:00
parent 41495b0c08
commit 7ad0cbea76
2 changed files with 65 additions and 167 deletions
+4 -1
View File
@@ -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))
+61 -166
View File
@@ -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:
"""构造合法日线 dfdate + 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 新 = 1032 个重复去重后保留 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 只 symbol2 只 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: 每只新增 3staging 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