Stream biz_state metric persist to cut peak memory on large BGP tables.

Write JSONL iteratively, flush via chunked PG execute_values, skip mega raw in DB, and yield bgp_route rows instead of building a full list.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-09-23 16:23:55 +08:00
parent 08231c779e
commit 9b6452971f
7 changed files with 423 additions and 136 deletions

View file

@ -2,9 +2,11 @@
from __future__ import annotations from __future__ import annotations
import json
import logging import logging
import re import re
import threading import threading
from collections.abc import Iterable
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
@ -152,16 +154,23 @@ def _resolve_collect_profile(profile_id: str):
return None return None
_METRIC_CHUNK = 2000
def _persist_lldp_rows( def _persist_lldp_rows(
db, db,
*, *,
batch: BizStateBatch, batch: BizStateBatch,
cmd_row: BizStateBatchCommand, cmd_row: BizStateBatchCommand,
records: list[dict[str, Any]], records: Iterable[dict[str, Any]],
) -> int: ) -> int:
"""Insert LLDP rows from an iterable (chunk-friendly)."""
n = 0 n = 0
seen: set[tuple[str, str, str]] = set() seen: set[tuple[str, str, str]] = set()
buf: list[BizStateLldpNeighbor] = []
for rec in records: for rec in records:
if not isinstance(rec, dict):
continue
local_if = str(rec.get("local_if") or "").strip()[:128] local_if = str(rec.get("local_if") or "").strip()[:128]
remote_sys = str(rec.get("remote_sys") or "").strip()[:256] remote_sys = str(rec.get("remote_sys") or "").strip()[:256]
remote_if = str(rec.get("remote_if") or "").strip()[:128] remote_if = str(rec.get("remote_if") or "").strip()[:128]
@ -171,7 +180,7 @@ def _persist_lldp_rows(
if key in seen: if key in seen:
continue continue
seen.add(key) seen.add(key)
db.add( buf.append(
BizStateLldpNeighbor( BizStateLldpNeighbor(
id=uuid4().hex, id=uuid4().hex,
batch_id=batch.id, batch_id=batch.id,
@ -187,6 +196,190 @@ def _persist_lldp_rows(
) )
) )
n += 1 n += 1
if len(buf) >= _METRIC_CHUNK:
db.add_all(buf)
db.flush()
buf.clear()
if buf:
db.add_all(buf)
db.flush()
return n
def _db_is_postgres(db) -> bool:
try:
bind = db.get_bind() if hasattr(db, "get_bind") else getattr(db, "bind", None)
name = str(getattr(getattr(bind, "dialect", None), "name", "") or "").lower()
return name in ("postgresql", "postgres")
except Exception:
return False
def _persist_metric_chunk_bulk(
db,
*,
batch: BizStateBatch,
cmd_row: BizStateBatchCommand,
metric_id: str,
chunk: list[dict[str, Any]],
seq_start: int,
) -> int:
if not chunk:
return 0
now = _utcnow()
buf = [
{
"id": uuid4().hex,
"batch_id": batch.id,
"batch_command_id": cmd_row.id,
"task_id": batch.task_id,
"ne_id": batch.ne_id,
"metric_id": metric_id,
"seq": seq_start + i,
"data_json": dict(rec),
"collected_at": now,
}
for i, rec in enumerate(chunk)
if isinstance(rec, dict) and rec
]
if not buf:
return 0
db.bulk_insert_mappings(BizStateMetricRow, buf)
db.flush()
return len(buf)
def _persist_metric_chunk_copy(
db,
*,
batch: BizStateBatch,
cmd_row: BizStateBatchCommand,
metric_id: str,
chunk: list[dict[str, Any]],
seq_start: int,
) -> int:
"""Postgres fast path via execute_values; falls back to bulk on error."""
if not chunk:
return 0
now = _utcnow()
rows: list[tuple[Any, ...]] = []
for i, rec in enumerate(chunk):
if not isinstance(rec, dict) or not rec:
continue
rows.append(
(
uuid4().hex,
batch.id,
cmd_row.id,
batch.task_id or "",
batch.ne_id or "",
metric_id,
seq_start + i,
json.dumps(rec, ensure_ascii=False, default=str, separators=(",", ":")),
now,
)
)
if not rows:
return 0
try:
from psycopg2.extras import execute_values # type: ignore
except ImportError:
return _persist_metric_chunk_bulk(
db,
batch=batch,
cmd_row=cmd_row,
metric_id=metric_id,
chunk=chunk,
seq_start=seq_start,
)
# Unwrap SQLAlchemy connection → DBAPI (psycopg2) connection
sa_conn = db.connection()
dbapi = sa_conn.connection
driver = getattr(dbapi, "dbapi_connection", None) or getattr(
dbapi, "driver_connection", None
) or dbapi
sql = (
"INSERT INTO biz_state_metric_row "
"(id,batch_id,batch_command_id,task_id,ne_id,metric_id,seq,data_json,collected_at) "
"VALUES %s"
)
with driver.cursor() as cur:
execute_values(
cur,
sql,
rows,
template="(%s,%s,%s,%s,%s,%s,%s,%s::jsonb,%s)",
page_size=len(rows),
)
db.flush()
return len(rows)
def _persist_metric_rows_from_spool(
db,
*,
batch: BizStateBatch,
cmd_row: BizStateBatchCommand,
metric_id: str,
records_rel_path: str,
) -> int:
"""Stream JSONL → DB in chunks (never loads full table into memory)."""
from .spool import iter_record_chunks
mid = str(metric_id or "").strip()
if not mid or not records_rel_path:
return 0
use_copy = _db_is_postgres(db)
n = 0
for chunk in iter_record_chunks(records_rel_path, chunk_size=_METRIC_CHUNK):
if use_copy:
try:
added = _persist_metric_chunk_copy(
db,
batch=batch,
cmd_row=cmd_row,
metric_id=mid,
chunk=chunk,
seq_start=n,
)
except Exception:
_log.exception("biz_state COPY failed; falling back to bulk")
use_copy = False
added = _persist_metric_chunk_bulk(
db,
batch=batch,
cmd_row=cmd_row,
metric_id=mid,
chunk=chunk,
seq_start=n,
)
else:
added = _persist_metric_chunk_bulk(
db,
batch=batch,
cmd_row=cmd_row,
metric_id=mid,
chunk=chunk,
seq_start=n,
)
n += added
return n
def _persist_lldp_rows_from_spool(
db,
*,
batch: BizStateBatch,
cmd_row: BizStateBatchCommand,
records_rel_path: str,
) -> int:
from .spool import iter_record_chunks
if not records_rel_path:
return 0
n = 0
for chunk in iter_record_chunks(records_rel_path, chunk_size=_METRIC_CHUNK):
n += _persist_lldp_rows(db, batch=batch, cmd_row=cmd_row, records=chunk)
return n return n
@ -216,7 +409,6 @@ _GENERIC_METRICS = {
"config_ospf", "config_ospf",
"config_isis", "config_isis",
} }
_METRIC_CHUNK = 2000
def _run_primary_parse_job(job: Any) -> tuple[bool, bool]: def _run_primary_parse_job(job: Any) -> tuple[bool, bool]:
@ -239,10 +431,11 @@ def _run_primary_parse_job(job: Any) -> tuple[bool, bool]:
m = re.search(r"(?i)total\s+number\s+of\s+routes\s*:\s*(\d+)", text) m = re.search(r"(?i)total\s+number\s+of\s+routes\s*:\s*(\d+)", text)
return int(m.group(1)) if m else 0 return int(m.group(1)) if m else 0
def _flush_item(item: SpooledCommand, *, records: list[dict[str, Any]] | None = None) -> None: def _flush_item(item: SpooledCommand, *, records: Any = None) -> None:
if records is not None and item.persist_kind: if records is not None and item.persist_kind:
item.records_rel_path = write_records(batch_id, item.id, records) rel, n = write_records(batch_id, item.id, records)
item.row_count = len(records) item.records_rel_path = rel
item.row_count = int(n or 0)
try: try:
write_meta(batch_id, item.id, item.to_meta()) write_meta(batch_id, item.id, item.to_meta())
except Exception: except Exception:
@ -400,21 +593,25 @@ def _run_primary_parse_job(job: Any) -> tuple[bool, bool]:
"enrich=" + ",".join(getattr(j, "from_aux", "") for j in job.enrich_joins) "enrich=" + ",".join(getattr(j, "from_aux", "") for j in job.enrich_joins)
) )
declared = int(primary.declared_total or 0) declared = int(primary.declared_total or 0)
nrec = len(records or [])
if declared > 0:
hints.append(f"declared={declared};parsed={nrec}")
if hints:
primary.message = ";".join(hints)[:1020]
primary.parse_status = "ok" primary.parse_status = "ok"
any_ok = True any_ok = True
persist_recs = None persist_recs = None
if job.metric_id == "lldp_neighbor": if job.metric_id == "lldp_neighbor":
primary.persist_kind = "lldp" primary.persist_kind = "lldp"
persist_recs = list(records or []) persist_recs = records
elif job.metric_id in _GENERIC_METRICS: elif job.metric_id in _GENERIC_METRICS:
primary.persist_kind = "metric" primary.persist_kind = "metric"
persist_recs = list(records or []) persist_recs = records
_flush_item(primary, records=persist_recs) _flush_item(primary, records=persist_recs)
nrec = int(primary.row_count or 0)
if declared > 0:
hints.append(f"declared={declared};parsed={nrec}")
if hints:
primary.message = ";".join(hints)[:1020]
try:
write_meta(batch_id, primary.id, primary.to_meta())
except Exception:
pass
_ = fsm_tables # kept for hints above _ = fsm_tables # kept for hints above
except Exception as exc: except Exception as exc:
any_fail = True any_fail = True
@ -450,7 +647,7 @@ def _flush_spooled_commands(
pending: list[Any], pending: list[Any],
) -> tuple[int, int]: ) -> tuple[int, int]:
"""Insert SpooledCommand rows (+ metric/lldp) in one transaction. Returns (cmds, rows).""" """Insert SpooledCommand rows (+ metric/lldp) in one transaction. Returns (cmds, rows)."""
from .spool import SpooledCommand, raw_max_bytes, read_raw_text, read_records from .spool import SpooledCommand, raw_max_bytes, read_raw_text, spool_file_size
if not pending: if not pending:
return 0, 0 return 0, 0
@ -467,25 +664,38 @@ def _flush_spooled_commands(
raw = "" raw = ""
truncated = False truncated = False
line_count = int(getattr(item, "raw_line_count", 0) or 0) line_count = int(getattr(item, "raw_line_count", 0) or 0)
raw_size = 0
if item.raw_rel_path: if item.raw_rel_path:
from .spool import count_file_lines from .spool import count_file_lines
raw_size = spool_file_size(item.raw_rel_path)
if line_count <= 0: if line_count <= 0:
try: try:
line_count = count_file_lines(item.raw_rel_path) line_count = count_file_lines(item.raw_rel_path)
except Exception: except Exception:
line_count = 0 line_count = 0
raw = read_raw_text(item.raw_rel_path, max_bytes=max_raw) # Mega outputs stay on spool only — do not load full CLI into Postgres.
if max_raw > 0 and "[truncated" in raw: if max_raw > 0 and raw_size > max_raw:
truncated = True truncated = True
# Prefer full-file line count; fall back to stored text. raw = (
if line_count <= 0 and raw: f"[raw_on_spool={item.raw_rel_path}; size={raw_size}B; "
f"cap={max_raw}B; omitted_from_db]\n"
)
else:
raw = read_raw_text(item.raw_rel_path, max_bytes=max_raw)
if max_raw > 0 and "[truncated" in raw:
truncated = True
if line_count <= 0 and raw and not truncated:
from .spool import count_text_lines from .spool import count_text_lines
line_count = count_text_lines(raw) line_count = count_text_lines(raw)
msg = str(item.message or "").strip() msg = str(item.message or "").strip()
if truncated: if truncated:
note = f"raw_truncated@{max_raw}B" note = (
f"raw_on_spool@{raw_size}B"
if raw_size > max_raw > 0
else f"raw_truncated@{max_raw}B"
)
msg = f"{msg}; {note}" if msg else note msg = f"{msg}; {note}" if msg else note
cmd_row = BizStateBatchCommand( cmd_row = BizStateBatchCommand(
id=item.id, id=item.id,
@ -506,26 +716,26 @@ def _flush_spooled_commands(
) )
db.add(cmd_row) db.add(cmd_row)
if item.persist_kind == "metric" and item.records_rel_path: if item.persist_kind == "metric" and item.records_rel_path:
records = read_records(item.records_rel_path)
mid = str(item.metric_id or "").strip() mid = str(item.metric_id or "").strip()
if mid and records: if mid:
n = _persist_metric_rows( n = _persist_metric_rows_from_spool(
db, db,
batch=batch, batch=batch,
cmd_row=cmd_row, cmd_row=cmd_row,
metric_id=mid, metric_id=mid,
records=records, records_rel_path=item.records_rel_path,
) )
cmd_row.row_count = n cmd_row.row_count = n
rows_n += n rows_n += n
elif item.persist_kind == "lldp" and item.records_rel_path: elif item.persist_kind == "lldp" and item.records_rel_path:
records = read_records(item.records_rel_path) n = _persist_lldp_rows_from_spool(
if records: db,
n = _persist_lldp_rows( batch=batch,
db, batch=batch, cmd_row=cmd_row, records=records cmd_row=cmd_row,
) records_rel_path=item.records_rel_path,
cmd_row.row_count = n )
rows_n += n cmd_row.row_count = n
rows_n += n
db.commit() db.commit()
return len(items), rows_n return len(items), rows_n
@ -541,45 +751,6 @@ def _flush_spooled_commands(
raise raise
def _persist_metric_rows(
db,
*,
batch: BizStateBatch,
cmd_row: BizStateBatchCommand,
metric_id: str,
records: list[dict[str, Any]],
) -> int:
"""Bulk-insert generic metric rows (JSON payload per row)."""
mid = str(metric_id or "").strip()
if not mid or not records:
return 0
buf: list[dict[str, Any]] = []
n = 0
for i, rec in enumerate(records):
if not isinstance(rec, dict) or not rec:
continue
buf.append(
{
"id": uuid4().hex,
"batch_id": batch.id,
"batch_command_id": cmd_row.id,
"task_id": batch.task_id,
"ne_id": batch.ne_id,
"metric_id": mid,
"seq": i,
"data_json": dict(rec),
"collected_at": _utcnow(),
}
)
n += 1
if len(buf) >= _METRIC_CHUNK:
db.bulk_insert_mappings(BizStateMetricRow, buf)
buf.clear()
if buf:
db.bulk_insert_mappings(BizStateMetricRow, buf)
return n
def _finish_task(task_id: str, *, error: str = "") -> None: def _finish_task(task_id: str, *, error: str = "") -> None:
db = SessionLocal() db = SessionLocal()
try: try:
@ -825,11 +996,12 @@ def _run_collect_lane(
pending = [] pending = []
persist.submit(batch_id, chunk) persist.submit(batch_id, chunk)
def _queue(item: SpooledCommand, *, records: list[dict[str, Any]] | None = None) -> None: def _queue(item: SpooledCommand, *, records: Any = None) -> None:
nonlocal cmd_count nonlocal cmd_count
if records is not None and item.persist_kind: if records is not None and item.persist_kind:
item.records_rel_path = write_records(batch_id, item.id, records) rel, n = write_records(batch_id, item.id, records)
item.row_count = len(records) item.records_rel_path = rel
item.row_count = int(n or 0)
try: try:
write_meta(batch_id, item.id, item.to_meta()) write_meta(batch_id, item.id, item.to_meta())
except Exception: except Exception:

View file

@ -341,7 +341,7 @@ def run_primary_with_bundle(
textfsm_command: str = "", textfsm_command: str = "",
params: dict[str, str] | None = None, params: dict[str, str] | None = None,
enrich_joins: list[EnrichJoin] | None = None, enrich_joins: list[EnrichJoin] | None = None,
) -> tuple[list[dict[str, Any]], dict[str, list[dict[str, Any]]], list[str]]: ) -> tuple[Any, dict[str, list[dict[str, Any]]], list[str]]:
records, fsm_tables, keys = run_parser( records, fsm_tables, keys = run_parser(
parser_id, parser_id,
raw_text=bundle.raws.get("primary") or "", raw_text=bundle.raws.get("primary") or "",
@ -356,5 +356,8 @@ def run_primary_with_bundle(
fsm_tables_extra=bundle.fsm_extra, fsm_tables_extra=bundle.fsm_extra,
) )
if enrich_joins: if enrich_joins:
# Enrich needs random access — materialize only when joins are declared.
if not isinstance(records, list):
records = list(records or [])
apply_enrich_joins(records, bundle.aux_records, enrich_joins) apply_enrich_joins(records, bundle.aux_records, enrich_joins)
return records, fsm_tables, keys return records, fsm_tables, keys

View file

@ -34,6 +34,7 @@ Complex joins that cannot be expressed as equal-field copy still go in
from __future__ import annotations from __future__ import annotations
import inspect import inspect
from collections.abc import Iterable
from typing import Any, Callable, Mapping, Sequence from typing import Any, Callable, Mapping, Sequence
from ...lldp_shared import resolve_vendor_key from ...lldp_shared import resolve_vendor_key
@ -231,6 +232,11 @@ def run_parser(
"params": params or {}, "params": params or {},
} }
records = fn(**filtered) records = fn(**filtered)
if not isinstance(records, list): # Allow list or streaming Iterable (generator); reject bare str/bytes.
if records is None:
records = []
elif isinstance(records, (str, bytes)):
records = []
elif not isinstance(records, list) and not isinstance(records, Iterable):
records = [] records = []
return records, fsm_tables, used_keys return records, fsm_tables, used_keys

View file

@ -205,7 +205,6 @@ def _skip_noise_line(line: str) -> bool:
def _emit_route( def _emit_route(
out: list[dict[str, Any]],
seen: set[str], seen: set[str],
*, *,
local_as: str, local_as: str,
@ -219,38 +218,36 @@ def _emit_route(
rest: str, rest: str,
flags: str, flags: str,
path_continuation: bool = False, path_continuation: bool = False,
) -> None: ) -> dict[str, Any] | None:
if not _looks_like_prefix(net): if not _looks_like_prefix(net):
return return None
if nh and not _looks_like_ip_or_prefix(nh): if nh and not _looks_like_ip_or_prefix(nh):
if re.search(r"[A-Za-z]", nh): if re.search(r"[A-Za-z]", nh):
return return None
nh = "" nh = ""
key = _route_dedupe_key(rd=rd, net=net, nh=nh or "") key = _route_dedupe_key(rd=rd, net=net, nh=nh or "")
if key in seen: if key in seen:
return return None
seen.add(key) seen.add(key)
metric, loc, tag, path = _split_rest(rest, path_continuation=path_continuation) metric, loc, tag, path = _split_rest(rest, path_continuation=path_continuation)
out.append( return {
{ "local_as": local_as[:16],
"local_as": local_as[:16], "afi": afi[:32],
"afi": afi[:32], "vrf": vrf[:128],
"vrf": vrf[:128], "neighbor": neighbor[:128],
"neighbor": neighbor[:128], "direction": direction[:8],
"direction": direction[:8], "rd": (rd or "")[:64],
"rd": (rd or "")[:64], "network": net[:128],
"network": net[:128], "next_hop": (nh or "")[:128],
"next_hop": (nh or "")[:128], "metric": metric[:32],
"metric": metric[:32], "loc_prf": loc[:32],
"loc_prf": loc[:32], "tag": tag[:32],
"tag": tag[:32], "path": path[:256],
"path": path[:256], "status_codes": re.sub(r"\s+", "", (flags or "").strip())[:16],
"status_codes": re.sub(r"\s+", "", (flags or "").strip())[:16], "as_num": "",
"as_num": "", "state": "",
"state": "", "pfx_rcd": "",
"pfx_rcd": "", }
}
)
def _hand_parse( def _hand_parse(
@ -262,9 +259,8 @@ def _hand_parse(
neighbor: str = "", neighbor: str = "",
direction: str = "", direction: str = "",
**_kw: Any, **_kw: Any,
) -> list[dict[str, Any]]: ):
"""Parse neighbor in/out tables; join heavy IPv6 / From / metric wraps.""" """Yield neighbor in/out route rows (streaming; joins heavy IPv6 / From wraps)."""
out: list[dict[str, Any]] = []
seen: set[str] = set() seen: set[str] = set()
pending_net = "" pending_net = ""
pending_flags = "" pending_flags = ""
@ -273,13 +269,12 @@ def _hand_parse(
current_rd = "" current_rd = ""
current_vrf = vrf current_vrf = vrf
def _flush_pending(*, rest: str = "", path_continuation: bool = False) -> None: def _flush_pending(*, rest: str = "", path_continuation: bool = False) -> dict[str, Any] | None:
nonlocal pending_net, pending_flags, pending_nh, pending_rest nonlocal pending_net, pending_flags, pending_nh, pending_rest
if not pending_net: if not pending_net:
return return None
merged = " ".join(x for x in (pending_rest, rest) if x).strip() merged = " ".join(x for x in (pending_rest, rest) if x).strip()
_emit_route( row = _emit_route(
out,
seen, seen,
local_as=local_as, local_as=local_as,
afi=afi, afi=afi,
@ -297,20 +292,24 @@ def _hand_parse(
pending_flags = "" pending_flags = ""
pending_nh = "" pending_nh = ""
pending_rest = "" pending_rest = ""
return row
def _start_pending(*, net: str, nh: str, flags: str, rest: str) -> None: def _start_pending(*, net: str, nh: str, flags: str, rest: str):
nonlocal pending_net, pending_flags, pending_nh, pending_rest nonlocal pending_net, pending_flags, pending_nh, pending_rest
_flush_pending() flushed = _flush_pending()
pending_net = net pending_net = net
pending_flags = flags pending_flags = flags
pending_nh = nh pending_nh = nh
pending_rest = rest pending_rest = rest
return flushed
for raw in _normalize_cli_text(raw_text).splitlines(): for raw in _normalize_cli_text(raw_text).splitlines():
line = raw.rstrip() line = raw.rstrip()
rd_m = _RD_RE.match(line.strip()) rd_m = _RD_RE.match(line.strip())
if rd_m: if rd_m:
_flush_pending() row = _flush_pending()
if row:
yield row
current_rd = (rd_m.group("rd") or "").strip() current_rd = (rd_m.group("rd") or "").strip()
vrf_from_rd = (rd_m.group("vrf") or "").strip() vrf_from_rd = (rd_m.group("vrf") or "").strip()
if vrf_from_rd: if vrf_from_rd:
@ -331,15 +330,21 @@ def _hand_parse(
pending_nh = first pending_nh = first
more = " ".join(parts[1:]) more = " ".join(parts[1:])
if more and _rest_looks_complete(more): if more and _rest_looks_complete(more):
_flush_pending(rest=more, path_continuation=True) row = _flush_pending(rest=more, path_continuation=True)
if row:
yield row
elif more: elif more:
pending_rest = " ".join(x for x in (pending_rest, more) if x) pending_rest = " ".join(x for x in (pending_rest, more) if x)
continue continue
if not _looks_like_prefix(first): if not _looks_like_prefix(first):
_flush_pending(rest=tok, path_continuation=True) row = _flush_pending(rest=tok, path_continuation=True)
if row:
yield row
continue continue
else: else:
_flush_pending(rest=tok, path_continuation=True) row = _flush_pending(rest=tok, path_continuation=True)
if row:
yield row
continue continue
# Peel optional status codes (* i / *i / > …) # Peel optional status codes (* i / *i / > …)
@ -355,9 +360,10 @@ def _hand_parse(
if m and _looks_like_prefix(m.group("net")) and _looks_like_ip_or_prefix(m.group("nh")): if m and _looks_like_prefix(m.group("net")) and _looks_like_ip_or_prefix(m.group("nh")):
rest = (m.group("rest") or "").strip() rest = (m.group("rest") or "").strip()
if _rest_looks_complete(rest): if _rest_looks_complete(rest):
_flush_pending() row = _flush_pending()
_emit_route( if row:
out, yield row
row = _emit_route(
seen, seen,
local_as=local_as, local_as=local_as,
afi=afi, afi=afi,
@ -370,19 +376,27 @@ def _hand_parse(
rest=rest, rest=rest,
flags=flags, flags=flags,
) )
if row:
yield row
else: else:
# Empty rest or From-only → wait for metric/path wrap row = _start_pending(
_start_pending(net=m.group("net"), nh=m.group("nh"), flags=flags, rest=rest) net=m.group("net"), nh=m.group("nh"), flags=flags, rest=rest
)
if row:
yield row
continue continue
# Network alone → wait for next-hop wrap # Network alone → wait for next-hop wrap
m_net = _NET_ONLY_RE.match(body) m_net = _NET_ONLY_RE.match(body)
if m_net and _looks_like_prefix(m_net.group("net")): if m_net and _looks_like_prefix(m_net.group("net")):
_start_pending(net=m_net.group("net"), nh="", flags=flags, rest="") row = _start_pending(net=m_net.group("net"), nh="", flags=flags, rest="")
if row:
yield row
continue continue
_flush_pending() row = _flush_pending()
return out if row:
yield row
def normalize_bgp_route( def normalize_bgp_route(
@ -391,8 +405,8 @@ def normalize_bgp_route(
command: str = "", command: str = "",
params: dict[str, str] | None = None, params: dict[str, str] | None = None,
**_kw: Any, **_kw: Any,
) -> list[dict[str, Any]]: ):
"""Hand-only normalize (see module docstring).""" """Hand-only normalize; returns [] or a streaming iterator of route dicts."""
raw_text = _normalize_cli_text(raw_text) raw_text = _normalize_cli_text(raw_text)
if _empty_if_total_zero(raw_text): if _empty_if_total_zero(raw_text):
return [] return []

View file

@ -6,6 +6,7 @@ import json
import logging import logging
import re import re
import shutil import shutil
from collections.abc import Iterable, Iterator, Mapping
from dataclasses import dataclass, field from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@ -78,15 +79,23 @@ def count_text_lines(text: str | None) -> int:
return s.count("\n") + (0 if s.endswith("\n") else 1) return s.count("\n") + (0 if s.endswith("\n") else 1)
def write_records(batch_id: str, cmd_id: str, records: list[dict[str, Any]]) -> str: def write_records(
"""Write parsed records as JSONL; return relative path.""" batch_id: str,
cmd_id: str,
records: Iterable[Mapping[str, Any]] | None,
) -> tuple[str, int]:
"""Stream parsed records as JSONL; return (relative path, row count)."""
_, _, rec_path = _cmd_paths(batch_id, cmd_id) _, _, rec_path = _cmd_paths(batch_id, cmd_id)
n = 0
with rec_path.open("w", encoding="utf-8", errors="replace") as fh: with rec_path.open("w", encoding="utf-8", errors="replace") as fh:
for rec in records or []: for rec in records or ():
fh.write(json.dumps(rec, ensure_ascii=False, default=str)) if not isinstance(rec, Mapping):
continue
fh.write(json.dumps(dict(rec), ensure_ascii=False, default=str, separators=(",", ":")))
fh.write("\n") fh.write("\n")
n += 1
rel = rec_path.resolve().relative_to(spool_root()) rel = rec_path.resolve().relative_to(spool_root())
return str(rel).replace("\\", "/") return str(rel).replace("\\", "/"), n
def write_meta(batch_id: str, cmd_id: str, meta: dict[str, Any]) -> str: def write_meta(batch_id: str, cmd_id: str, meta: dict[str, Any]) -> str:
@ -114,13 +123,26 @@ def read_raw_text(rel_path: str, *, max_bytes: int = 0) -> str:
return data.decode("utf-8", errors="replace") return data.decode("utf-8", errors="replace")
def read_records(rel_path: str) -> list[dict[str, Any]]: def spool_file_size(rel_path: str) -> int:
"""Byte size of a spool file; 0 if missing."""
if not rel_path: if not rel_path:
return [] return 0
path = (spool_root() / str(rel_path)).resolve() path = (spool_root() / str(rel_path)).resolve()
if not str(path).startswith(str(spool_root())) or not path.is_file(): if not str(path).startswith(str(spool_root())) or not path.is_file():
return [] return 0
out: list[dict[str, Any]] = [] try:
return int(path.stat().st_size)
except OSError:
return 0
def iter_records(rel_path: str) -> Iterator[dict[str, Any]]:
"""Yield one record dict at a time from JSONL (never loads full file)."""
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
with path.open("r", encoding="utf-8", errors="replace") as fh: with path.open("r", encoding="utf-8", errors="replace") as fh:
for line in fh: for line in fh:
line = line.strip() line = line.strip()
@ -131,8 +153,27 @@ def read_records(rel_path: str) -> list[dict[str, Any]]:
except json.JSONDecodeError: except json.JSONDecodeError:
continue continue
if isinstance(rec, dict): if isinstance(rec, dict):
out.append(rec) yield rec
return out
def iter_record_chunks(
rel_path: str, *, chunk_size: int = 2000
) -> Iterator[list[dict[str, Any]]]:
"""Yield lists of up to ``chunk_size`` records from JSONL."""
size = max(1, int(chunk_size or 2000))
buf: list[dict[str, Any]] = []
for rec in iter_records(rel_path):
buf.append(rec)
if len(buf) >= size:
yield buf
buf = []
if buf:
yield buf
def read_records(rel_path: str) -> list[dict[str, Any]]:
"""Load all records (tests / small payloads only — prefer iter_record_chunks)."""
return list(iter_records(rel_path))
@dataclass @dataclass

View file

@ -17,6 +17,7 @@ from netx_api.biz_state import spool as spool_mod
from netx_api.biz_state.spool import ( from netx_api.biz_state.spool import (
SpooledCommand, SpooledCommand,
clear_batch_spool, clear_batch_spool,
iter_record_chunks,
read_raw_text, read_raw_text,
read_records, read_records,
write_raw_text, write_raw_text,
@ -43,11 +44,23 @@ class BizStateSpoolIoTests(unittest.TestCase):
rel = write_raw_text(bid, cid, "show arp\nA B C") rel = write_raw_text(bid, cid, "show arp\nA B C")
self.assertTrue(rel.endswith("cmd1.raw.txt")) self.assertTrue(rel.endswith("cmd1.raw.txt"))
self.assertEqual(read_raw_text(rel), "show arp\nA B C") 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"}]) rrel, n = write_records(bid, cid, [{"ip": "1.1.1.1"}, {"ip": "2.2.2.2"}])
self.assertEqual(n, 2)
recs = read_records(rrel) recs = read_records(rrel)
self.assertEqual(len(recs), 2) self.assertEqual(len(recs), 2)
self.assertEqual(recs[0]["ip"], "1.1.1.1") self.assertEqual(recs[0]["ip"], "1.1.1.1")
def test_write_records_streams_generator(self) -> None:
def _gen():
for i in range(5):
yield {"i": i}
rrel, n = write_records("batch-g", "cmd-g", _gen())
self.assertEqual(n, 5)
chunks = list(iter_record_chunks(rrel, chunk_size=2))
self.assertEqual([len(c) for c in chunks], [2, 2, 1])
self.assertEqual(chunks[0][0]["i"], 0)
def test_raw_max_bytes_truncate(self) -> None: def test_raw_max_bytes_truncate(self) -> None:
bid = "b2" bid = "b2"
cid = "c2" cid = "c2"
@ -116,11 +129,12 @@ class BizStateFlushSpoolTests(unittest.TestCase):
def test_flush_inserts_command_and_metric_rows(self) -> None: def test_flush_inserts_command_and_metric_rows(self) -> None:
cid = uuid4().hex cid = uuid4().hex
raw_rel = write_raw_text("b-spool", cid, "ARP OUTPUT") raw_rel = write_raw_text("b-spool", cid, "ARP OUTPUT")
rec_rel = write_records( rec_rel, rec_n = write_records(
"b-spool", "b-spool",
cid, cid,
[{"ip": "10.0.0.1", "mac": "aaaa"}, {"ip": "10.0.0.2", "mac": "bbbb"}], [{"ip": "10.0.0.1", "mac": "aaaa"}, {"ip": "10.0.0.2", "mac": "bbbb"}],
) )
self.assertEqual(rec_n, 2)
pending = [ pending = [
SpooledCommand( SpooledCommand(
id=cid, id=cid,
@ -134,6 +148,7 @@ class BizStateFlushSpoolTests(unittest.TestCase):
message="spooled", message="spooled",
raw_rel_path=raw_rel, raw_rel_path=raw_rel,
records_rel_path=rec_rel, records_rel_path=rec_rel,
row_count=rec_n,
persist_kind="metric", persist_kind="metric",
) )
] ]
@ -158,6 +173,36 @@ class BizStateFlushSpoolTests(unittest.TestCase):
self.assertEqual(batch.command_count, 1) self.assertEqual(batch.command_count, 1)
self.assertEqual(batch.row_count, 2) self.assertEqual(batch.row_count, 2)
def test_flush_skips_mega_raw_into_db(self) -> None:
cid = uuid4().hex
big = "X" * (9 * 1024 * 1024)
raw_rel = write_raw_text("b-spool", cid, big)
rec_rel, rec_n = write_records("b-spool", cid, [{"k": 1}])
pending = [
SpooledCommand(
id=cid,
batch_id="b-spool",
metric_id="arp",
raw_command="show arp",
parse_status="ok",
message="ok",
raw_rel_path=raw_rel,
records_rel_path=rec_rel,
row_count=rec_n,
persist_kind="metric",
)
]
with patch.object(spool_mod.settings, "biz_state_raw_max_bytes", 8 * 1024 * 1024):
cmds, rows = runner._flush_spooled_commands("b-spool", pending)
self.assertEqual(cmds, 1)
self.assertEqual(rows, 1)
self.db.expire_all()
cmd = self.db.get(BizStateBatchCommand, cid)
assert cmd is not None
self.assertIn("raw_on_spool", cmd.raw_text)
self.assertNotIn("XXXX", cmd.raw_text)
self.assertLess(len(cmd.raw_text or ""), 500)
def test_flush_batches_multiple_without_per_cmd_sessions(self) -> None: def test_flush_batches_multiple_without_per_cmd_sessions(self) -> None:
pending: list[SpooledCommand] = [] pending: list[SpooledCommand] = []
for i in range(5): for i in range(5):

View file

@ -12,7 +12,7 @@ from netx_api.biz_state.enrich import apply_enrich_joins
from netx_api.biz_state.parsers.common.vrf_list import normalize_vrf_list from netx_api.biz_state.parsers.common.vrf_list import normalize_vrf_list
from netx_api.biz_state.parsers.zte import ( from netx_api.biz_state.parsers.zte import (
normalize_bgp_peer, normalize_bgp_peer,
normalize_bgp_route, normalize_bgp_route as _normalize_bgp_route_stream,
normalize_ip_route, normalize_ip_route,
normalize_ipv6_route, normalize_ipv6_route,
normalize_l2vpn_mac, normalize_l2vpn_mac,
@ -26,6 +26,12 @@ from netx_api.biz_state.profiles import AuxCommand, get_profile, metric_field_ma
from netx_api.ntc_parse import apply_rule from netx_api.ntc_parse import apply_rule
def normalize_bgp_route(**kwargs):
"""Materialize streaming parser for assertions / enrich joins."""
out = _normalize_bgp_route_stream(**kwargs)
return out if isinstance(out, list) else list(out)
_VRF_SAMPLE = """ _VRF_SAMPLE = """
Name Default RD Protocols VRF ID Name Default RD Protocols VRF ID
CUST_A 100:1 ipv4 1 CUST_A 100:1 ipv4 1