From f94a1455873cb1e93398dc8d5710dcf156e09f72 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 28 Jul 2026 20:36:11 +0800 Subject: [PATCH] =?UTF-8?q?fix(data):=20bs=5Feod=20=E5=8D=A1=E6=AD=BB?= =?UTF-8?q?=E6=A0=B9=E6=B2=BB=20=E2=80=94=20per-stock=20commit=20+=20baost?= =?UTF-8?q?ock=20=E8=B6=85=E6=97=B6=E5=8C=85=E8=A3=85=20+=20=E5=91=A8?= =?UTF-8?q?=E6=9C=9F=20relogin?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 根因(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。 --- scripts/data_platform/bs_eod.py | 106 +++- tests/data_platform/test_bs_eod_resilience.py | 453 ++++++++++++++++++ 2 files changed, 541 insertions(+), 18 deletions(-) create mode 100644 tests/data_platform/test_bs_eod_resilience.py diff --git a/scripts/data_platform/bs_eod.py b/scripts/data_platform/bs_eod.py index 00a194d..ae8bcf6 100644 --- a/scripts/data_platform/bs_eod.py +++ b/scripts/data_platform/bs_eod.py @@ -17,6 +17,7 @@ import logging import os import socket import sys +import threading import time from pathlib import Path @@ -38,6 +39,9 @@ DB = BASE / "data" / "quant_trading.db" VAL_DIR = BASE / "data" / "valuation_baostock" LOOKBACK = int(os.environ.get("LOOKBACK_DAYS", "7")) DAILY_LIMIT = int(os.environ.get("BS_DAILY_LIMIT", "48000")) +# ③ 周期 relogin: 每 N 只主动 relogin, 防服务端长连接 idle 超时 → CLOSE_WAIT 静默 hang. +# CLOSE_WAIT 是静默阻塞不抛异常, 被动 relogin (error_code != 0 触发) 救不了, 必须周期主动刷连接. +RELOGIN_EVERY = int(os.environ.get("BS_RELOGIN_EVERY", "500")) BS_INTERVAL = 0.3 QUERY_COUNT = 0 EXC_MAP = {"sh": "SSE", "sz": "SZSE"} @@ -114,6 +118,57 @@ def fetch_k(bs_code, fields, freq, start, end): return rows +# ======================== 超时包装 ======================== +# 根因: socket.setdefaulttimeout(30) 对 baostock 客户端的 rs.next() / recv 不可靠遵守, +# 服务端 hiccup 会无限阻塞主循环. 用 daemon 工作线程 + join(timeout) 强制上限. +# +# 线程方案 vs subprocess 隔离: 选线程. +# 理由: (1) baostock 模块级 singleton, 主线程 join 等子线程 → 单线程串行调用安全; +# (2) 超时后子线程 daemon 化泄漏, 后续 relogin 的 logout 关旧 socket → 旧线程 +# recv 报错自死, 新 login 走新 socket 不受污染 (relogin retry 兜底恢复); +# (3) subprocess 方案需重新 login (~1s) + 序列化 rows 复杂, 不值; +# (4) 与 akshare_static_download.call_ak_with_timeout 同范式. + +def _with_timeout(fn, args=(), kwargs=None, timeout=60): + """daemon 线程跑 fn(*args, **kwargs), timeout 秒未完成 raise TimeoutError. + + 超时后工作线程泄漏 (daemon=True, 进程退出时强杀); 主线程立即返回让上层 relogin. + 子线程异常透传给主线程 (BaseException 也捕获, 避免 daemon 吞 KeyboardInterrupt). + """ + if kwargs is None: + kwargs = {} + box = {"val": None, "exc": None} + + def worker(): + try: + box["val"] = fn(*args, **kwargs) + except BaseException as e: # noqa: BLE001 - 透传所有异常含 KeyboardInterrupt + box["exc"] = e + + t = threading.Thread(target=worker, daemon=True) + t.start() + t.join(timeout) + if t.is_alive(): + raise TimeoutError(f"{getattr(fn, '__name__', repr(fn))} 超过 {timeout}s") + if box["exc"] is not None: + raise box["exc"] + return box["val"] + + +def fetch_k_with_timeout(bs_code, fields, freq, start, end, timeout=60): + """fetch_k + 超时保护 (默认 60s; baostock hiccup 不再无限阻塞).""" + return _with_timeout( + fetch_k, + args=(bs_code, fields, freq, start, end), + timeout=timeout, + ) + + +def fetch_all_stocks_with_timeout(timeout=120): + """fetch_all_stocks + 超时保护 (全 A 列表一次性返回, 给 120s).""" + return _with_timeout(fetch_all_stocks, timeout=timeout) + + def upsert_daily(conn, code, prefix, rows): """日线 rows -> dbbardata('d') + valuation_baostock 当年 parquet 追加。""" if not rows: @@ -191,6 +246,27 @@ def upsert_15m(conn, code, prefix, rows): return len(db) +def _process_one_stock(conn, code, prefix, args, start, end): + """单只股票: fetch_k + upsert, 在 with conn 短事务里执行 (大事务根治). + + 每只股票一个事务 — hang/kill/异常最多丢 1 只, 已 commit 的其他股不受影响. + fetch 也包在事务里 (task 要求: "fetch_k + upsert 包在自己事务里"); + fetch 用 fetch_k_with_timeout 保护, 网络挂最多锁 timeout 秒. + 成功返 (n1_daily, n2_15m); 异常时 with conn 自动 ROLLBACK 该股, 异常上抛. + """ + bs_code = f"{prefix}.{code}" + n1 = 0 + n2 = 0 + with conn: # 显式短事务: 成功 commit / 异常 rollback (per-stock 原子) + if not args.no_daily: + d_rows = fetch_k_with_timeout(bs_code, DAILY_FIELDS, "d", start, end) + n1 = upsert_daily(conn, code, prefix, d_rows) + if not args.no_15m: + m_rows = fetch_k_with_timeout(bs_code, M15_FIELDS, "15", start, end) + n2 = upsert_15m(conn, code, prefix, m_rows) + return n1, n2 + + def main(): global QUERY_COUNT ap = argparse.ArgumentParser() @@ -210,7 +286,7 @@ def main(): sys.exit(2) try: - stocks = fetch_all_stocks() + stocks = fetch_all_stocks_with_timeout() except Exception as e: log.error("[FATAL] fetch_all: %s", e) sys.exit(1) @@ -227,45 +303,39 @@ def main(): stats = {"ok": 0, "empty": 0, "failed": 0, "db_rows": 0} limit_reached = False t0 = time.time() - conn.execute("BEGIN") + # 大事务根治: 不再 BEGIN/COMMIT 包全程. 每只股票 with conn 短事务独立提交, + # hang/kill/崩溃最多丢 1 只 (per-stock 隔离), 已 commit 的进度不丢. try: for i, (code, prefix) in enumerate(stocks): if QUERY_COUNT >= DAILY_LIMIT: log.warning("query %d 达防线 %d, graceful stop", QUERY_COUNT, DAILY_LIMIT) limit_reached = True break - bs_code = f"{prefix}.{code}" try: - n1 = 0 - if not args.no_daily: - d_rows = fetch_k(bs_code, DAILY_FIELDS, "d", start, end) - n1 = upsert_daily(conn, code, prefix, d_rows) - n2 = 0 - if not args.no_15m: - m_rows = fetch_k(bs_code, M15_FIELDS, "15", start, end) - n2 = upsert_15m(conn, code, prefix, m_rows) + n1, n2 = _process_one_stock(conn, code, prefix, args, start, end) stats["db_rows"] += n1 + n2 if n1 + n2: stats["ok"] += 1 else: stats["empty"] += 1 except Exception as e: + # _process_one_stock 异常已 rollback 该股, 其他股已 commit 不受影响 stats["failed"] += 1 - if stats["failed"] <= 5: + if stats["failed"] <= 5 or stats["failed"] % 100 == 0: log.warning("%s err: %s", code, e) if not relogin(): log.error("%s relogin 失败, 跳过", code) - if (i + 1) % 500 == 0: + if (i + 1) % RELOGIN_EVERY == 0: log.info("进度 %d/%d ok=%d empty=%d failed=%d q=%d (%.0fs)", i + 1, len(stocks), stats["ok"], stats["empty"], stats["failed"], QUERY_COUNT, time.time() - t0) + # ③ 主动 relogin (双保险之治本): 服务端长连接 idle 超时 → CLOSE_WAIT 静默 hang, + # 被动 relogin 不触发 (不抛异常), 必须周期主动 logout+login 刷新连接. + # ② _with_timeout 是治标兜底, 真挂了能打破; 这里治本避免走到那一步. + if not relogin(): + log.warning("周期 relogin 失败, 继续跑 (下次 fetch 失败时被动 relogin 兜底)") if i < len(stocks) - 1: time.sleep(BS_INTERVAL) - conn.execute("COMMIT") - except Exception as e: - conn.execute("ROLLBACK") - log.error("[FATAL] rollback: %s", e) - sys.exit(1) finally: conn.close() try: diff --git a/tests/data_platform/test_bs_eod_resilience.py b/tests/data_platform/test_bs_eod_resilience.py new file mode 100644 index 0000000..77e2337 --- /dev/null +++ b/tests/data_platform/test_bs_eod_resilience.py @@ -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/.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, 没到周期