Harden long biz_state collects: release DB during CLI and spool before flush.

Avoid idle Session disconnect after heavy timeout, and separate SSH collect from batched Postgres persist via disk spool.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-09-22 16:24:46 +08:00
parent c15ea00a23
commit 0eb65a839e
5 changed files with 1293 additions and 433 deletions

File diff suppressed because it is too large Load diff

160
netx_api/biz_state/spool.py Normal file
View file

@ -0,0 +1,160 @@
"""Filesystem spool for biz_state collect: CLI/raw + parsed records before DB flush."""
from __future__ import annotations
import json
import logging
import re
import shutil
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from ..config import settings
_log = logging.getLogger("netx.biz_state.spool")
_SAFE_RE = re.compile(r"[^A-Za-z0-9._-]+")
def spool_root() -> Path:
root = Path(str(getattr(settings, "biz_state_spool_dir", None) or "data/biz_state_spool"))
root.mkdir(parents=True, exist_ok=True)
return root.resolve()
def batch_spool_dir(batch_id: str) -> Path:
bid = _SAFE_RE.sub("_", str(batch_id or "").strip())[:64] or "unknown"
path = spool_root() / bid
path.mkdir(parents=True, exist_ok=True)
return path
def clear_batch_spool(batch_id: str) -> None:
bid = _SAFE_RE.sub("_", str(batch_id or "").strip())[:64]
if not bid:
return
path = (spool_root() / bid).resolve()
root = spool_root()
if not str(path).startswith(str(root)) or path == root:
return
if path.is_dir():
shutil.rmtree(path, ignore_errors=True)
def _cmd_paths(batch_id: str, cmd_id: str) -> tuple[Path, Path, Path]:
base = batch_spool_dir(batch_id)
cid = _SAFE_RE.sub("_", str(cmd_id or "").strip())[:64] or "cmd"
return base / f"{cid}.raw.txt", base / f"{cid}.meta.json", base / f"{cid}.records.jsonl"
def write_raw_text(batch_id: str, cmd_id: str, text: str) -> str:
"""Write CLI output; return path relative to spool root (posix)."""
raw_path, _, _ = _cmd_paths(batch_id, cmd_id)
raw_path.write_bytes(str(text or "").encode("utf-8", errors="replace"))
rel = raw_path.resolve().relative_to(spool_root())
return str(rel).replace("\\", "/")
def write_records(batch_id: str, cmd_id: str, records: list[dict[str, Any]]) -> str:
"""Write parsed records as JSONL; return relative path."""
_, _, rec_path = _cmd_paths(batch_id, cmd_id)
with rec_path.open("w", encoding="utf-8", errors="replace") as fh:
for rec in records or []:
fh.write(json.dumps(rec, ensure_ascii=False, default=str))
fh.write("\n")
rel = rec_path.resolve().relative_to(spool_root())
return str(rel).replace("\\", "/")
def write_meta(batch_id: str, cmd_id: str, meta: dict[str, Any]) -> str:
_, meta_path, _ = _cmd_paths(batch_id, cmd_id)
meta_path.write_text(
json.dumps(meta, ensure_ascii=False, default=str),
encoding="utf-8",
errors="replace",
)
rel = meta_path.resolve().relative_to(spool_root())
return str(rel).replace("\\", "/")
def read_raw_text(rel_path: str, *, max_bytes: int = 0) -> str:
if not rel_path:
return ""
path = (spool_root() / str(rel_path)).resolve()
if not str(path).startswith(str(spool_root())) or not path.is_file():
return ""
data = path.read_bytes()
cap = int(max_bytes or 0)
if cap > 0 and len(data) > cap:
text = data[:cap].decode("utf-8", errors="replace")
return text + f"\n...[truncated {cap} bytes cap]\n"
return data.decode("utf-8", errors="replace")
def read_records(rel_path: str) -> list[dict[str, Any]]:
if not rel_path:
return []
path = (spool_root() / str(rel_path)).resolve()
if not str(path).startswith(str(spool_root())) or not path.is_file():
return []
out: list[dict[str, Any]] = []
with path.open("r", encoding="utf-8", errors="replace") as fh:
for line in fh:
line = line.strip()
if not line:
continue
try:
rec = json.loads(line)
except json.JSONDecodeError:
continue
if isinstance(rec, dict):
out.append(rec)
return out
@dataclass
class SpooledCommand:
"""One command (primary or aux) collected on disk, awaiting DB flush."""
id: str
batch_id: str
task_item_id: str = ""
profile_id: str = ""
parser_id: str = ""
metric_id: str = ""
raw_command: str = ""
params_json: dict[str, Any] = field(default_factory=dict)
parse_status: str = ""
message: str = ""
raw_rel_path: str = ""
records_rel_path: str = ""
row_count: int = 0
# "" | "metric" | "lldp"
persist_kind: str = ""
def to_meta(self) -> dict[str, Any]:
return {
"id": self.id,
"batch_id": self.batch_id,
"task_item_id": self.task_item_id,
"profile_id": self.profile_id,
"parser_id": self.parser_id,
"metric_id": self.metric_id,
"raw_command": self.raw_command,
"params_json": dict(self.params_json or {}),
"parse_status": self.parse_status,
"message": self.message,
"raw_rel_path": self.raw_rel_path,
"records_rel_path": self.records_rel_path,
"row_count": self.row_count,
"persist_kind": self.persist_kind,
}
def persist_every_cmds() -> int:
return max(1, int(getattr(settings, "biz_state_persist_every_cmds", 8) or 8))
def raw_max_bytes() -> int:
return max(0, int(getattr(settings, "biz_state_raw_max_bytes", 8 * 1024 * 1024) or 0))

View file

@ -122,6 +122,11 @@ class Settings(BaseSettings):
biz_state_heavy_read_timeout_sec: int = 1500 biz_state_heavy_read_timeout_sec: int = 1500
biz_state_heavy_run_timeout_cap_sec: int = 2400 biz_state_heavy_run_timeout_cap_sec: int = 2400
biz_state_heavy_workers: int = 4 biz_state_heavy_workers: int = 4
# Collect spool: CLI/raw + parsed records on disk; flush to DB every N cmds.
biz_state_spool_dir: str = "data/biz_state_spool"
biz_state_persist_every_cmds: int = 8
# Cap raw_text loaded into Postgres from spool (0 = unlimited).
biz_state_raw_max_bytes: int = 8 * 1024 * 1024
# Managed NE exec: max CLI commands per request (lab can raise; hard-capped in ne_exec). # Managed NE exec: max CLI commands per request (lab can raise; hard-capped in ne_exec).
ne_exec_max_commands: int = 5 ne_exec_max_commands: int = 5
# Opt-in: allow per-NE exec_policy (linux_shell/unrestricted). Default off — UI hidden. # Opt-in: allow per-NE exec_policy (linux_shell/unrestricted). Default off — UI hidden.

View file

@ -0,0 +1,232 @@
"""Finalize / reconnect after long-collect DB disconnect."""
from __future__ import annotations
import unittest
from unittest.mock import patch
from sqlalchemy import create_engine
from sqlalchemy.exc import OperationalError
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
from netx_api.biz_state import collect_runner as runner
from netx_api.db import Base
from netx_api.models import BizStateBatch, BizStateTask
class BizStateCollectFinalizeTests(unittest.TestCase):
def setUp(self) -> None:
# StaticPool: all sessions share one :memory: SQLite DB.
engine = create_engine(
"sqlite+pysqlite:///:memory:",
future=True,
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
TestingSession = sessionmaker(
bind=engine, autoflush=False, autocommit=False, expire_on_commit=False
)
Base.metadata.create_all(bind=engine)
self.Session = TestingSession
self.db = TestingSession()
self.task = BizStateTask(
id="t-finalize",
source="managed",
ne_id="ne1",
ne_name="PE1",
status="running",
collect_running=True,
interval_sec=300,
)
self.db.add(self.task)
self.batch = BizStateBatch(
id="b-finalize",
task_id="t-finalize",
status="running",
command_count=55,
row_count=6606,
message="",
)
self.db.add(self.batch)
self.db.commit()
def tearDown(self) -> None:
self.db.close()
def test_finalize_partial_after_heavy_timeout(self) -> None:
with patch.object(runner, "SessionLocal", self.Session):
with patch(
"netx_api.biz_state.compare_service.try_auto_compare_for_task",
return_value=None,
):
status = runner._finalize_batch_status(
batch_id="b-finalize",
task_id="t-finalize",
cmd_count=40,
total_rows=1000,
any_fail=True,
any_ok=True,
lane_errors=["RuntimeError: biz_state_heavy_timeout (2400s)"],
)
self.assertEqual(status, "partial")
self.db.expire_all()
batch = self.db.get(BizStateBatch, "b-finalize")
assert batch is not None
self.assertEqual(batch.status, "partial")
self.assertEqual(batch.command_count, 55)
self.assertEqual(batch.row_count, 6606)
self.assertIn("biz_state_heavy_timeout", batch.message or "")
self.assertIsNotNone(batch.ended_at)
def test_finalize_retries_once_on_operational_error(self) -> None:
calls = {"n": 0}
real_session = self.Session
class FlakySession:
"""First commit raises OperationalError; subsequent sessions work."""
def __init__(self) -> None:
self._inner = real_session()
self._failed = False
def get(self, *args, **kwargs):
return self._inner.get(*args, **kwargs)
def commit(self) -> None:
calls["n"] += 1
if calls["n"] == 1:
self._failed = True
raise OperationalError(
"UPDATE",
{},
Exception(
"consuming input failed: server closed the connection unexpectedly"
),
)
self._inner.commit()
def rollback(self) -> None:
try:
self._inner.rollback()
except Exception:
pass
def close(self) -> None:
self._inner.close()
def connection(self):
# Avoid invalidate() wiping StaticPool's only connection in tests.
raise RuntimeError("skip invalidate in test")
def add(self, *args, **kwargs):
return self._inner.add(*args, **kwargs)
def session_factory():
return FlakySession()
with patch.object(runner, "SessionLocal", session_factory):
with patch(
"netx_api.biz_state.compare_service.try_auto_compare_for_task",
return_value=None,
):
status = runner._finalize_batch_status(
batch_id="b-finalize",
task_id="t-finalize",
cmd_count=55,
total_rows=6606,
any_fail=True,
any_ok=True,
lane_errors=["RuntimeError: biz_state_heavy_timeout (2400s)"],
)
self.assertEqual(status, "partial")
self.assertEqual(calls["n"], 2)
self.db.expire_all()
batch = self.db.get(BizStateBatch, "b-finalize")
assert batch is not None
self.assertEqual(batch.status, "partial")
self.assertIn("biz_state_heavy_timeout", batch.message or "")
def test_fail_batch_retries_on_stale_connection(self) -> None:
calls = {"n": 0}
real_session = self.Session
class FlakySession:
def __init__(self) -> None:
self._inner = real_session()
def get(self, *args, **kwargs):
return self._inner.get(*args, **kwargs)
def commit(self) -> None:
calls["n"] += 1
if calls["n"] == 1:
raise OperationalError(
"UPDATE",
{},
Exception("server closed the connection unexpectedly"),
)
self._inner.commit()
def rollback(self) -> None:
try:
self._inner.rollback()
except Exception:
pass
def close(self) -> None:
self._inner.close()
def connection(self):
raise RuntimeError("skip invalidate in test")
with patch.object(runner, "SessionLocal", FlakySession):
runner._fail_batch_status("b-finalize", "RuntimeError: boom")
self.assertEqual(calls["n"], 2)
self.db.expire_all()
batch = self.db.get(BizStateBatch, "b-finalize")
assert batch is not None
self.assertEqual(batch.status, "failed")
self.assertIn("boom", batch.message or "")
def test_run_db_with_reconnect_reraises_non_stale(self) -> None:
calls = {"n": 0}
class BoomSession:
def commit(self) -> None:
pass
def rollback(self) -> None:
pass
def close(self) -> None:
pass
def connection(self):
raise RuntimeError("no conn")
def fn(db) -> None:
calls["n"] += 1
raise ValueError("not a disconnect")
with patch.object(runner, "SessionLocal", BoomSession):
with self.assertRaises(ValueError):
runner._run_db_with_reconnect(fn, label="test")
self.assertEqual(calls["n"], 1)
def test_is_stale_db_connection(self) -> None:
self.assertTrue(
runner._is_stale_db_connection(
OperationalError("x", {}, Exception("server closed the connection"))
)
)
self.assertTrue(
runner._is_stale_db_connection(
RuntimeError("consuming input failed: server closed the connection unexpectedly")
)
)
self.assertFalse(runner._is_stale_db_connection(ValueError("nope")))
if __name__ == "__main__":
unittest.main()

View file

@ -0,0 +1,189 @@
"""biz_state collect spool: disk collect + batched DB flush."""
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from uuid import uuid4
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
from netx_api.biz_state import collect_runner as runner
from netx_api.biz_state import spool as spool_mod
from netx_api.biz_state.spool import (
SpooledCommand,
clear_batch_spool,
read_raw_text,
read_records,
write_raw_text,
write_records,
)
from netx_api.db import Base
from netx_api.models import BizStateBatch, BizStateBatchCommand, BizStateMetricRow, BizStateTask
class BizStateSpoolIoTests(unittest.TestCase):
def setUp(self) -> None:
self._tmpdir = tempfile.TemporaryDirectory()
self.root = Path(self._tmpdir.name)
self._patcher = patch.object(spool_mod.settings, "biz_state_spool_dir", str(self.root))
self._patcher.start()
def tearDown(self) -> None:
self._patcher.stop()
self._tmpdir.cleanup()
def test_write_read_raw_and_records(self) -> None:
bid = "batch1"
cid = "cmd1"
rel = write_raw_text(bid, cid, "show arp\nA B C")
self.assertTrue(rel.endswith("cmd1.raw.txt"))
self.assertEqual(read_raw_text(rel), "show arp\nA B C")
rrel = write_records(bid, cid, [{"ip": "1.1.1.1"}, {"ip": "2.2.2.2"}])
recs = read_records(rrel)
self.assertEqual(len(recs), 2)
self.assertEqual(recs[0]["ip"], "1.1.1.1")
def test_raw_max_bytes_truncate(self) -> None:
bid = "b2"
cid = "c2"
rel = write_raw_text(bid, cid, "x" * 100)
text = read_raw_text(rel, max_bytes=20)
self.assertIn("truncated", text)
self.assertLess(len(text), 80)
def test_clear_batch_spool(self) -> None:
bid = "b3"
write_raw_text(bid, "c", "hi")
self.assertTrue((self.root / "b3").is_dir())
clear_batch_spool(bid)
self.assertFalse((self.root / "b3").exists())
class BizStateFlushSpoolTests(unittest.TestCase):
def setUp(self) -> None:
self._tmpdir = tempfile.TemporaryDirectory()
self.root = Path(self._tmpdir.name)
self._spool_patch = patch.object(
spool_mod.settings, "biz_state_spool_dir", str(self.root)
)
self._spool_patch.start()
engine = create_engine(
"sqlite+pysqlite:///:memory:",
future=True,
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
TestingSession = sessionmaker(
bind=engine, autoflush=False, autocommit=False, expire_on_commit=False
)
Base.metadata.create_all(bind=engine)
self.Session = TestingSession
self.db = TestingSession()
self.task = BizStateTask(
id="t-spool",
source="managed",
ne_id="ne1",
ne_name="PE1",
status="running",
collect_running=True,
interval_sec=300,
)
self.db.add(self.task)
self.batch = BizStateBatch(
id="b-spool",
task_id="t-spool",
status="running",
command_count=0,
row_count=0,
)
self.db.add(self.batch)
self.db.commit()
self._session_patch = patch.object(runner, "SessionLocal", TestingSession)
self._session_patch.start()
def tearDown(self) -> None:
self._session_patch.stop()
self._spool_patch.stop()
self.db.close()
self._tmpdir.cleanup()
def test_flush_inserts_command_and_metric_rows(self) -> None:
cid = uuid4().hex
raw_rel = write_raw_text("b-spool", cid, "ARP OUTPUT")
rec_rel = write_records(
"b-spool",
cid,
[{"ip": "10.0.0.1", "mac": "aaaa"}, {"ip": "10.0.0.2", "mac": "bbbb"}],
)
pending = [
SpooledCommand(
id=cid,
batch_id="b-spool",
task_item_id="item1",
profile_id="zte.arp",
parser_id="zte_arp",
metric_id="arp",
raw_command="show arp",
parse_status="ok",
message="spooled",
raw_rel_path=raw_rel,
records_rel_path=rec_rel,
persist_kind="metric",
)
]
cmds, rows = runner._flush_spooled_commands("b-spool", pending)
self.assertEqual(cmds, 1)
self.assertEqual(rows, 2)
self.assertEqual(pending, [])
self.db.expire_all()
cmd = self.db.get(BizStateBatchCommand, cid)
assert cmd is not None
self.assertEqual(cmd.parse_status, "ok")
self.assertEqual(cmd.raw_text, "ARP OUTPUT")
self.assertEqual(cmd.row_count, 2)
n = (
self.db.query(BizStateMetricRow)
.filter(BizStateMetricRow.batch_command_id == cid)
.count()
)
self.assertEqual(n, 2)
batch = self.db.get(BizStateBatch, "b-spool")
assert batch is not None
self.assertEqual(batch.command_count, 1)
self.assertEqual(batch.row_count, 2)
def test_flush_batches_multiple_without_per_cmd_sessions(self) -> None:
pending: list[SpooledCommand] = []
for i in range(5):
cid = uuid4().hex
raw_rel = write_raw_text("b-spool", cid, f"out-{i}")
pending.append(
SpooledCommand(
id=cid,
batch_id="b-spool",
raw_command=f"show x {i}",
parse_status="skipped_custom",
message="custom_raw",
raw_rel_path=raw_rel,
)
)
cmds, rows = runner._flush_spooled_commands("b-spool", pending)
self.assertEqual(cmds, 5)
self.assertEqual(rows, 0)
self.db.expire_all()
n = (
self.db.query(BizStateBatchCommand)
.filter(BizStateBatchCommand.batch_id == "b-spool")
.count()
)
self.assertEqual(n, 5)
if __name__ == "__main__":
unittest.main()