fix(data): bs_eod 卡死根治 — per-stock commit + baostock 超时包装 + 周期 relogin
根因(py-spy dump + netstat CLOSE_WAIT 实证): baostock 服务端关长连接→CLOSE_WAIT, send_msg 静默阻塞不抛异常, socket.setdefaulttimeout 不被 baostock 自己 socket 遵守, relogin 只在 error_code≠0 救不了; 一把大事务全程持 WAL 锁阻断全库。修复: per-stock commit 去大事务 + _with_timeout 线程超时包 fetch_k 打破静默 hang + 周期 relogin 每500主动刷连接。VPS --limit 3 验证 11s 不 hang。
This commit is contained in:
@@ -0,0 +1,453 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""TDD for bs_eod.py resilience fixes (大事务根治 + baostock 超时包装).
|
||||
|
||||
测试覆盖:
|
||||
1. _with_timeout: 成功返回 / 慢调用 TimeoutError / 子异常透传 / kwargs 传递
|
||||
2. fetch_k_with_timeout: 包装 fetch_k 并应用超时
|
||||
3. _process_one_stock: 每只股票独立事务 (3rd 失败 → 前 2 已提交)
|
||||
4. per-stock 隔离: 单只股票部分失败 → 整股 rollback, 不影响其他股
|
||||
5. upsert_daily / upsert_15m 不回归
|
||||
|
||||
Mac 无 baostock 也能跑 (mock baostock 模块 + patch fetch_k 不依赖真实 baostock).
|
||||
"""
|
||||
import sqlite3
|
||||
import sys
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
# Mock baostock before import (Mac 可能没装 / 不依赖网络)
|
||||
if "baostock" not in sys.modules:
|
||||
sys.modules["baostock"] = MagicMock()
|
||||
|
||||
from scripts.data_platform import bs_eod # noqa: E402
|
||||
|
||||
|
||||
# ---------- Fixtures ----------
|
||||
|
||||
@pytest.fixture
|
||||
def tmp_db(tmp_path):
|
||||
"""临时 sqlite DB 带 dbbardata 表 (复用生产 schema 主键)."""
|
||||
db_path = tmp_path / "test.db"
|
||||
conn = sqlite3.connect(str(db_path))
|
||||
conn.execute(
|
||||
"CREATE TABLE dbbardata ("
|
||||
"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, "
|
||||
"PRIMARY KEY (symbol, exchange, datetime, interval))"
|
||||
)
|
||||
conn.commit()
|
||||
yield conn
|
||||
conn.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tmp_valdir(tmp_path, monkeypatch):
|
||||
"""重定向 VAL_DIR 到 tmp_path (避免污染 Mac 当前年度 parquet)."""
|
||||
val = tmp_path / "valuation"
|
||||
val.mkdir(parents=True, exist_ok=True)
|
||||
monkeypatch.setattr(bs_eod, "VAL_DIR", val)
|
||||
return val
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def reset_query_count():
|
||||
"""每个测试前重置 QUERY_COUNT (全局可变状态)."""
|
||||
original = bs_eod.QUERY_COUNT
|
||||
bs_eod.QUERY_COUNT = 0
|
||||
yield
|
||||
bs_eod.QUERY_COUNT = original
|
||||
|
||||
|
||||
# ---------- _with_timeout 单测 ----------
|
||||
|
||||
def test_with_timeout_returns_value_on_success():
|
||||
"""快速函数应在超时前返回值."""
|
||||
def fast_fn(x, y):
|
||||
return x + y
|
||||
out = bs_eod._with_timeout(fast_fn, args=(2, 3), timeout=5)
|
||||
assert out == 5
|
||||
|
||||
|
||||
def test_with_timeout_raises_timeouterror_on_hang():
|
||||
"""hang 函数应在 timeout 内 raise TimeoutError, 不无限阻塞."""
|
||||
def slow_fn():
|
||||
time.sleep(10)
|
||||
return "done"
|
||||
t0 = time.time()
|
||||
with pytest.raises(TimeoutError):
|
||||
bs_eod._with_timeout(slow_fn, timeout=1)
|
||||
elapsed = time.time() - t0
|
||||
# 必须在 ~1s 内返回, 不能等 sleep(10) — 这就是 fix 的核心
|
||||
assert elapsed < 3, f"timeout 不生效, elapsed={elapsed:.1f}s"
|
||||
|
||||
|
||||
def test_with_timeout_propagates_exception():
|
||||
"""子线程异常应透传给主线程."""
|
||||
def boom():
|
||||
raise ValueError("kaboom")
|
||||
with pytest.raises(ValueError, match="kaboom"):
|
||||
bs_eod._with_timeout(boom, timeout=5)
|
||||
|
||||
|
||||
def test_with_timeout_passes_kwargs():
|
||||
"""kwargs 正确传给 fn."""
|
||||
def fn(a, b, c=99):
|
||||
return (a, b, c)
|
||||
out = bs_eod._with_timeout(fn, args=(1, 2), kwargs={"c": 3}, timeout=5)
|
||||
assert out == (1, 2, 3)
|
||||
|
||||
|
||||
def test_with_timeout_breaks_close_wait_silent_hang():
|
||||
"""针对 netstat 实测 CLOSE_WAIT 静默 hang 根因 (VPS 现场证据).
|
||||
|
||||
CLOSE_WAIT 半关闭态: baostock 客户端 recv 不干净处理 EOF, 卡在内部循环 —
|
||||
**不抛异常, 不触发 error_code != 0**, 所以:
|
||||
- socket.setdefaulttimeout(30) 救不了 (客户端不死在 socket 层)
|
||||
- 被动 relogin (只在异常时触发) 永远不触发
|
||||
唯一能打破静默 hang 的就是外层 _with_timeout. 此处显式模拟"无异常纯阻塞"调用.
|
||||
"""
|
||||
def close_wait_silent_block():
|
||||
# 模拟 baostock 客户端在 CLOSE_WAIT 上死循环: 永不返回, 也不抛异常
|
||||
while True:
|
||||
time.sleep(0.1)
|
||||
|
||||
t0 = time.time()
|
||||
with pytest.raises(TimeoutError):
|
||||
bs_eod._with_timeout(close_wait_silent_block, timeout=1)
|
||||
elapsed = time.time() - t0
|
||||
assert elapsed < 3, f"CLOSE_WAIT 静默 hang 未被外层超时打破: elapsed={elapsed:.1f}s"
|
||||
|
||||
|
||||
# ---------- fetch_k_with_timeout / fetch_all_stocks_with_timeout ----------
|
||||
|
||||
def test_fetch_k_with_timeout_delegates_to_fetch_k():
|
||||
"""fetch_k_with_timeout 应透传参数给 fetch_k 并返回其结果."""
|
||||
sentinel = [("row",)]
|
||||
with patch.object(bs_eod, "fetch_k", return_value=sentinel) as m:
|
||||
out = bs_eod.fetch_k_with_timeout(
|
||||
"sh.600000", "fields", "d", "2026-01-01", "2026-01-31"
|
||||
)
|
||||
assert out is sentinel
|
||||
m.assert_called_once_with("sh.600000", "fields", "d", "2026-01-01", "2026-01-31")
|
||||
|
||||
|
||||
def test_fetch_k_with_timeout_raises_when_fetch_k_hangs():
|
||||
"""hang 的 fetch_k 应被超时杀掉 (baostock 服务端 hiccup 不再无限阻塞)."""
|
||||
def hang(*a, **kw):
|
||||
time.sleep(10)
|
||||
with patch.object(bs_eod, "fetch_k", side_effect=hang):
|
||||
with pytest.raises(TimeoutError):
|
||||
bs_eod.fetch_k_with_timeout(
|
||||
"sh.600000", "f", "d", "s", "e", timeout=1
|
||||
)
|
||||
|
||||
|
||||
def test_fetch_all_stocks_with_timeout_delegates():
|
||||
"""fetch_all_stocks_with_timeout 应委托给 fetch_all_stocks."""
|
||||
sentinel = [("000001", "sz")]
|
||||
with patch.object(bs_eod, "fetch_all_stocks", return_value=sentinel) as m:
|
||||
out = bs_eod.fetch_all_stocks_with_timeout()
|
||||
assert out is sentinel
|
||||
m.assert_called_once_with()
|
||||
|
||||
|
||||
# ---------- _process_one_stock (per-stock 短事务) ----------
|
||||
|
||||
def _make_daily_rows(code, n=2):
|
||||
"""模拟 baostock daily 返回的 row (15 列 DAILY_FIELDS)."""
|
||||
out = []
|
||||
for i in range(n):
|
||||
date = f"2026-07-{i + 1:02d}"
|
||||
out.append([date, code, "10", "11", "9", "10.5", "1000", "10000",
|
||||
"1.5", "0.5", "20", "5", "1", "2", "0"])
|
||||
return out
|
||||
|
||||
|
||||
def _make_15m_rows(code, n=1):
|
||||
"""模拟 baostock 15min 返回的 row (9 列 M15_FIELDS)."""
|
||||
out = []
|
||||
for i in range(n):
|
||||
out.append(["2026-07-01", "20260701094500000", code,
|
||||
"10", "11", "9", "10.5", "1000", "10000"])
|
||||
return out
|
||||
|
||||
|
||||
def test_process_one_stock_commits_each_stock_independently(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""3 只股票, 第 3 只 fetch_k 抛错 → 前 2 只数据已 commit, 第 3 只 rollback 不影响.
|
||||
|
||||
这是大事务根治的核心断言: 旧版 BEGIN/COMMIT 包全程, 第 3 只失败会 rollback 前 2 只;
|
||||
新版每只独立 with conn, 前 2 只已落盘.
|
||||
"""
|
||||
stocks = [("600001", "sh"), ("600002", "sh"), ("600003", "sh")]
|
||||
|
||||
def fake_fetch(bs_code, fields, freq, start, end):
|
||||
if bs_code.endswith("600003"):
|
||||
raise RuntimeError("simulated baostock error")
|
||||
return _make_daily_rows(bs_code.split(".")[1], n=2)
|
||||
|
||||
args = MagicMock(no_daily=False, no_15m=True)
|
||||
with patch.object(bs_eod, "fetch_k_with_timeout", side_effect=fake_fetch):
|
||||
ok, fail = 0, 0
|
||||
for code, prefix in stocks:
|
||||
try:
|
||||
bs_eod._process_one_stock(
|
||||
tmp_db, code, prefix, args, "2026-07-01", "2026-07-31"
|
||||
)
|
||||
ok += 1
|
||||
except Exception:
|
||||
fail += 1
|
||||
|
||||
assert ok == 2
|
||||
assert fail == 1
|
||||
|
||||
# 前 2 只的行已持久化 (即使第 3 只失败也不会 rollback)
|
||||
rows = tmp_db.execute(
|
||||
"SELECT DISTINCT symbol FROM dbbardata ORDER BY symbol"
|
||||
).fetchall()
|
||||
assert [r[0] for r in rows] == ["600001", "600002"]
|
||||
|
||||
|
||||
def test_process_one_stock_daily_and_15m(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""单只股票同时跑 daily + 15min, 两表都写入."""
|
||||
def fake_fetch(bs_code, fields, freq, start, end):
|
||||
code = bs_code.split(".")[1]
|
||||
if freq == "d":
|
||||
return _make_daily_rows(code, 1)
|
||||
if freq == "15":
|
||||
return _make_15m_rows(code, 1)
|
||||
return []
|
||||
|
||||
args = MagicMock(no_daily=False, no_15m=False)
|
||||
with patch.object(bs_eod, "fetch_k_with_timeout", side_effect=fake_fetch):
|
||||
n1, n2 = bs_eod._process_one_stock(
|
||||
tmp_db, "000001", "sz", args, "2026-07-01", "2026-07-31"
|
||||
)
|
||||
assert n1 == 1
|
||||
assert n2 == 1
|
||||
|
||||
rows = tmp_db.execute(
|
||||
"SELECT interval, COUNT(*) FROM dbbardata GROUP BY interval"
|
||||
).fetchall()
|
||||
by_interval = dict(rows)
|
||||
assert by_interval.get("d") == 1
|
||||
assert by_interval.get("15m") == 1
|
||||
|
||||
|
||||
def test_process_one_stock_no_daily_skips_daily(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""--no-daily 跳过日线, 只跑 15min."""
|
||||
def fake_fetch(bs_code, fields, freq, start, end):
|
||||
if freq == "d":
|
||||
pytest.fail("daily 不应被调用 (--no-daily)")
|
||||
return _make_15m_rows(bs_code.split(".")[1], 1)
|
||||
|
||||
args = MagicMock(no_daily=True, no_15m=False)
|
||||
with patch.object(bs_eod, "fetch_k_with_timeout", side_effect=fake_fetch):
|
||||
n1, n2 = bs_eod._process_one_stock(
|
||||
tmp_db, "000001", "sz", args, "2026-07-01", "2026-07-31"
|
||||
)
|
||||
assert n1 == 0
|
||||
assert n2 == 1
|
||||
rows = tmp_db.execute("SELECT interval FROM dbbardata").fetchall()
|
||||
assert rows == [("15m",)]
|
||||
|
||||
|
||||
def test_process_one_stock_no_15m_skips_15m(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""--no-15m 跳过 15min, 只跑 daily."""
|
||||
def fake_fetch(bs_code, fields, freq, start, end):
|
||||
if freq == "15":
|
||||
pytest.fail("15min 不应被调用 (--no-15m)")
|
||||
return _make_daily_rows(bs_code.split(".")[1], 1)
|
||||
|
||||
args = MagicMock(no_daily=False, no_15m=True)
|
||||
with patch.object(bs_eod, "fetch_k_with_timeout", side_effect=fake_fetch):
|
||||
n1, n2 = bs_eod._process_one_stock(
|
||||
tmp_db, "000001", "sz", args, "2026-07-01", "2026-07-31"
|
||||
)
|
||||
assert n1 == 1
|
||||
assert n2 == 0
|
||||
|
||||
|
||||
def test_process_one_stock_atomic_per_stock(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""单只股票部分失败 → 整股 rollback (per-stock 原子性).
|
||||
|
||||
daily 成功, 15min 抛错 → 整只股票事务 rollback, daily 行也不留.
|
||||
与 test_process_one_stock_commits_each_stock_independently 共同证明:
|
||||
隔离边界是"股票"而非"全程".
|
||||
"""
|
||||
def fake_fetch(bs_code, fields, freq, start, end):
|
||||
code = bs_code.split(".")[1]
|
||||
if freq == "d":
|
||||
return _make_daily_rows(code, 1)
|
||||
if freq == "15":
|
||||
raise RuntimeError("15min failed")
|
||||
return []
|
||||
|
||||
args = MagicMock(no_daily=False, no_15m=False)
|
||||
with patch.object(bs_eod, "fetch_k_with_timeout", side_effect=fake_fetch):
|
||||
with pytest.raises(RuntimeError, match="15min failed"):
|
||||
bs_eod._process_one_stock(
|
||||
tmp_db, "000001", "sz", args, "2026-07-01", "2026-07-31"
|
||||
)
|
||||
|
||||
count = tmp_db.execute("SELECT COUNT(*) FROM dbbardata").fetchone()[0]
|
||||
assert count == 0, "单只股票部分失败应整体 rollback"
|
||||
|
||||
|
||||
def test_process_one_stock_persists_after_reopen(
|
||||
tmp_db, tmp_valdir, reset_query_count, tmp_path
|
||||
):
|
||||
"""持久化保证: with conn commit 后, 重开 DB 数据仍在 (模拟 kill 后恢复)."""
|
||||
def fake_fetch(bs_code, fields, freq, start, end):
|
||||
return _make_daily_rows(bs_code.split(".")[1], 1)
|
||||
|
||||
args = MagicMock(no_daily=False, no_15m=True)
|
||||
with patch.object(bs_eod, "fetch_k_with_timeout", side_effect=fake_fetch):
|
||||
bs_eod._process_one_stock(
|
||||
tmp_db, "000001", "sz", args, "2026-07-01", "2026-07-31"
|
||||
)
|
||||
|
||||
# 关闭并重开 DB, 验证 commit 持久
|
||||
db_path = tmp_db.execute("PRAGMA database_list").fetchall()[0][2]
|
||||
tmp_db.close()
|
||||
conn2 = sqlite3.connect(db_path)
|
||||
try:
|
||||
count = conn2.execute("SELECT COUNT(*) FROM dbbardata").fetchone()[0]
|
||||
assert count == 1
|
||||
finally:
|
||||
conn2.close()
|
||||
|
||||
|
||||
# ---------- upsert_daily / upsert_15m 不回归 ----------
|
||||
|
||||
def test_upsert_daily_writes_dbbardata_and_valuation_parquet(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""upsert_daily: rows -> dbbardata('d') + valuation_baostock/<year>.parquet."""
|
||||
rows = _make_daily_rows("600000", 1)
|
||||
n = bs_eod.upsert_daily(tmp_db, "600000", "sh", rows)
|
||||
tmp_db.commit()
|
||||
assert n == 1
|
||||
|
||||
db_row = tmp_db.execute(
|
||||
"SELECT symbol, exchange, datetime, interval, close_price FROM dbbardata"
|
||||
).fetchone()
|
||||
assert db_row[0] == "600000"
|
||||
assert db_row[1] == "SSE"
|
||||
assert db_row[3] == "d"
|
||||
assert db_row[4] == 10.5
|
||||
|
||||
yr = pd.Timestamp.now().year
|
||||
parquet_path = tmp_valdir / f"{yr}.parquet"
|
||||
assert parquet_path.exists()
|
||||
df = pd.read_parquet(parquet_path)
|
||||
assert "symbol" in df.columns
|
||||
assert (df["symbol"] == "600000").any()
|
||||
|
||||
|
||||
def test_upsert_15m_writes_dbbardata(
|
||||
tmp_db, tmp_valdir, reset_query_count
|
||||
):
|
||||
"""upsert_15m: rows -> dbbardata('15m') 正确字段 + datetime 拼接."""
|
||||
rows = _make_15m_rows("000001", 1)
|
||||
n = bs_eod.upsert_15m(tmp_db, "000001", "sz", rows)
|
||||
tmp_db.commit()
|
||||
assert n == 1
|
||||
|
||||
row = tmp_db.execute(
|
||||
"SELECT symbol, exchange, datetime, interval, close_price FROM dbbardata"
|
||||
).fetchone()
|
||||
assert row[0] == "000001"
|
||||
assert row[1] == "SZSE"
|
||||
assert row[2] == "2026-07-01 09:45:00"
|
||||
assert row[3] == "15m"
|
||||
assert row[4] == 10.5
|
||||
|
||||
|
||||
def test_upsert_daily_empty_rows_no_op(tmp_db, tmp_valdir):
|
||||
"""空 rows 不写."""
|
||||
n = bs_eod.upsert_daily(tmp_db, "600000", "sh", [])
|
||||
assert n == 0
|
||||
count = tmp_db.execute("SELECT COUNT(*) FROM dbbardata").fetchone()[0]
|
||||
assert count == 0
|
||||
|
||||
|
||||
def test_upsert_15m_empty_rows_no_op(tmp_db, tmp_valdir):
|
||||
"""空 rows 不写."""
|
||||
n = bs_eod.upsert_15m(tmp_db, "000001", "sz", [])
|
||||
assert n == 0
|
||||
count = tmp_db.execute("SELECT COUNT(*) FROM dbbardata").fetchone()[0]
|
||||
assert count == 0
|
||||
|
||||
|
||||
# ---------- ③ 周期 relogin (CLOSE_WAIT 治本之主动防御) ----------
|
||||
|
||||
@pytest.fixture
|
||||
def isolated_main_env(tmp_path, monkeypatch, reset_query_count):
|
||||
"""main() 集成测试环境: 重定向 DB/VAL_DIR/sys.argv, mock baostock 模块."""
|
||||
monkeypatch.setattr(bs_eod, "DB", tmp_path / "fake.db")
|
||||
val = tmp_path / "val"
|
||||
val.mkdir(parents=True, exist_ok=True)
|
||||
monkeypatch.setattr(bs_eod, "VAL_DIR", val)
|
||||
monkeypatch.setattr(sys, "argv", ["bs_eod.py"])
|
||||
# bs 模块 mock: login 成功 (error_code="0"), logout 不抛
|
||||
mock_bs = MagicMock()
|
||||
mock_bs.login.return_value.error_code = "0"
|
||||
monkeypatch.setattr(bs_eod, "bs", mock_bs)
|
||||
# 避免 BS_INTERVAL sleep 拖慢测试
|
||||
monkeypatch.setattr(bs_eod, "BS_INTERVAL", 0.0)
|
||||
return tmp_path
|
||||
|
||||
|
||||
def test_periodic_relogin_called_every_n_stocks(isolated_main_env, monkeypatch):
|
||||
"""③ 周期 relogin: 每 RELOGIN_EVERY 只主动 relogin, 主动刷 baostock 连接防 CLOSE_WAIT."""
|
||||
monkeypatch.setattr(bs_eod, "RELOGIN_EVERY", 2)
|
||||
stocks = [(f"60000{i}", "sh") for i in range(6)] # 6 只 / 每 2 只 → 3 次
|
||||
with patch.object(bs_eod, "fetch_all_stocks_with_timeout", return_value=stocks), \
|
||||
patch.object(bs_eod, "_process_one_stock", return_value=(1, 0)), \
|
||||
patch.object(bs_eod, "relogin", return_value=True) as m_rel:
|
||||
with pytest.raises(SystemExit) as exc:
|
||||
bs_eod.main()
|
||||
assert exc.value.code == 0
|
||||
# i+1=2,4,6 三次进度块, 每次都调 relogin → 3 次
|
||||
assert m_rel.call_count == 3
|
||||
|
||||
|
||||
def test_periodic_relogin_failure_does_not_crash_main(isolated_main_env, monkeypatch):
|
||||
"""③ 周期 relogin 失败 (返 False) 不应中断主循环, 后续 fetch 失败时被动 relogin 兜底."""
|
||||
monkeypatch.setattr(bs_eod, "RELOGIN_EVERY", 2)
|
||||
stocks = [(f"60000{i}", "sh") for i in range(4)]
|
||||
with patch.object(bs_eod, "fetch_all_stocks_with_timeout", return_value=stocks), \
|
||||
patch.object(bs_eod, "_process_one_stock", return_value=(1, 0)) as m_proc, \
|
||||
patch.object(bs_eod, "relogin", return_value=False) as m_rel:
|
||||
with pytest.raises(SystemExit) as exc:
|
||||
bs_eod.main()
|
||||
# 4 只都跑了 (周期 relogin 返 False 仅 warning, 不 crash)
|
||||
assert exc.value.code == 0
|
||||
assert m_proc.call_count == 4
|
||||
assert m_rel.call_count == 2 # i+1=2,4
|
||||
|
||||
|
||||
def test_periodic_relogin_disabled_when_relogin_every_huge(isolated_main_env, monkeypatch):
|
||||
"""RELOGIN_EVERY 极大时, 主循环不触发周期 relogin (回归保护)."""
|
||||
monkeypatch.setattr(bs_eod, "RELOGIN_EVERY", 10000)
|
||||
stocks = [(f"60000{i}", "sh") for i in range(5)]
|
||||
with patch.object(bs_eod, "fetch_all_stocks_with_timeout", return_value=stocks), \
|
||||
patch.object(bs_eod, "_process_one_stock", return_value=(1, 0)), \
|
||||
patch.object(bs_eod, "relogin", return_value=True) as m_rel:
|
||||
with pytest.raises(SystemExit):
|
||||
bs_eod.main()
|
||||
assert m_rel.call_count == 0 # 5 < 10000, 没到周期
|
||||
Reference in New Issue
Block a user