diff --git a/netx_api/biz_state/collect_runner.py b/netx_api/biz_state/collect_runner.py index d2ac5be..2636826 100644 --- a/netx_api/biz_state/collect_runner.py +++ b/netx_api/biz_state/collect_runner.py @@ -2,9 +2,11 @@ from __future__ import annotations +import json import logging import re import threading +from collections.abc import Iterable from concurrent.futures import ThreadPoolExecutor from datetime import datetime from typing import Any @@ -152,16 +154,23 @@ def _resolve_collect_profile(profile_id: str): return None +_METRIC_CHUNK = 2000 + + def _persist_lldp_rows( db, *, batch: BizStateBatch, cmd_row: BizStateBatchCommand, - records: list[dict[str, Any]], + records: Iterable[dict[str, Any]], ) -> int: + """Insert LLDP rows from an iterable (chunk-friendly).""" n = 0 seen: set[tuple[str, str, str]] = set() + buf: list[BizStateLldpNeighbor] = [] for rec in records: + if not isinstance(rec, dict): + continue local_if = str(rec.get("local_if") or "").strip()[:128] remote_sys = str(rec.get("remote_sys") or "").strip()[:256] remote_if = str(rec.get("remote_if") or "").strip()[:128] @@ -171,7 +180,7 @@ def _persist_lldp_rows( if key in seen: continue seen.add(key) - db.add( + buf.append( BizStateLldpNeighbor( id=uuid4().hex, batch_id=batch.id, @@ -187,6 +196,190 @@ def _persist_lldp_rows( ) ) 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 @@ -216,7 +409,6 @@ _GENERIC_METRICS = { "config_ospf", "config_isis", } -_METRIC_CHUNK = 2000 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) 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: - item.records_rel_path = write_records(batch_id, item.id, records) - item.row_count = len(records) + rel, n = write_records(batch_id, item.id, records) + item.records_rel_path = rel + item.row_count = int(n or 0) try: write_meta(batch_id, item.id, item.to_meta()) 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) ) 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" any_ok = True persist_recs = None if job.metric_id == "lldp_neighbor": primary.persist_kind = "lldp" - persist_recs = list(records or []) + persist_recs = records elif job.metric_id in _GENERIC_METRICS: primary.persist_kind = "metric" - persist_recs = list(records or []) + persist_recs = records _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 except Exception as exc: any_fail = True @@ -450,7 +647,7 @@ def _flush_spooled_commands( pending: list[Any], ) -> tuple[int, int]: """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: return 0, 0 @@ -467,25 +664,38 @@ def _flush_spooled_commands( raw = "" truncated = False line_count = int(getattr(item, "raw_line_count", 0) or 0) + raw_size = 0 if item.raw_rel_path: from .spool import count_file_lines + raw_size = spool_file_size(item.raw_rel_path) if line_count <= 0: try: line_count = count_file_lines(item.raw_rel_path) except Exception: line_count = 0 - raw = read_raw_text(item.raw_rel_path, max_bytes=max_raw) - if max_raw > 0 and "[truncated" in raw: + # Mega outputs stay on spool only — do not load full CLI into Postgres. + if max_raw > 0 and raw_size > max_raw: truncated = True - # Prefer full-file line count; fall back to stored text. - if line_count <= 0 and raw: + 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 line_count = count_text_lines(raw) msg = str(item.message or "").strip() 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 cmd_row = BizStateBatchCommand( id=item.id, @@ -506,26 +716,26 @@ def _flush_spooled_commands( ) db.add(cmd_row) if item.persist_kind == "metric" and item.records_rel_path: - records = read_records(item.records_rel_path) mid = str(item.metric_id or "").strip() - if mid and records: - n = _persist_metric_rows( + if mid: + n = _persist_metric_rows_from_spool( db, batch=batch, cmd_row=cmd_row, metric_id=mid, - records=records, + records_rel_path=item.records_rel_path, ) cmd_row.row_count = n rows_n += n elif item.persist_kind == "lldp" and item.records_rel_path: - records = read_records(item.records_rel_path) - if records: - n = _persist_lldp_rows( - db, batch=batch, cmd_row=cmd_row, records=records - ) - cmd_row.row_count = n - rows_n += n + n = _persist_lldp_rows_from_spool( + db, + batch=batch, + cmd_row=cmd_row, + records_rel_path=item.records_rel_path, + ) + cmd_row.row_count = n + rows_n += n db.commit() return len(items), rows_n @@ -541,45 +751,6 @@ def _flush_spooled_commands( 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: db = SessionLocal() try: @@ -825,11 +996,12 @@ def _run_collect_lane( pending = [] 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 if records is not None and item.persist_kind: - item.records_rel_path = write_records(batch_id, item.id, records) - item.row_count = len(records) + rel, n = write_records(batch_id, item.id, records) + item.records_rel_path = rel + item.row_count = int(n or 0) try: write_meta(batch_id, item.id, item.to_meta()) except Exception: diff --git a/netx_api/biz_state/collect_session.py b/netx_api/biz_state/collect_session.py index 17d0b3a..2c1f863 100644 --- a/netx_api/biz_state/collect_session.py +++ b/netx_api/biz_state/collect_session.py @@ -341,7 +341,7 @@ def run_primary_with_bundle( textfsm_command: str = "", params: dict[str, str] | 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( parser_id, raw_text=bundle.raws.get("primary") or "", @@ -356,5 +356,8 @@ def run_primary_with_bundle( fsm_tables_extra=bundle.fsm_extra, ) 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) return records, fsm_tables, keys diff --git a/netx_api/biz_state/parsers/__init__.py b/netx_api/biz_state/parsers/__init__.py index 939b5ea..fa92855 100644 --- a/netx_api/biz_state/parsers/__init__.py +++ b/netx_api/biz_state/parsers/__init__.py @@ -34,6 +34,7 @@ Complex joins that cannot be expressed as equal-field copy still go in from __future__ import annotations import inspect +from collections.abc import Iterable from typing import Any, Callable, Mapping, Sequence from ...lldp_shared import resolve_vendor_key @@ -231,6 +232,11 @@ def run_parser( "params": params or {}, } 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 = [] return records, fsm_tables, used_keys diff --git a/netx_api/biz_state/parsers/zte/bgp_route.py b/netx_api/biz_state/parsers/zte/bgp_route.py index fac0c15..cdb04d2 100644 --- a/netx_api/biz_state/parsers/zte/bgp_route.py +++ b/netx_api/biz_state/parsers/zte/bgp_route.py @@ -205,7 +205,6 @@ def _skip_noise_line(line: str) -> bool: def _emit_route( - out: list[dict[str, Any]], seen: set[str], *, local_as: str, @@ -219,38 +218,36 @@ def _emit_route( rest: str, flags: str, path_continuation: bool = False, -) -> None: +) -> dict[str, Any] | None: if not _looks_like_prefix(net): - return + return None if nh and not _looks_like_ip_or_prefix(nh): if re.search(r"[A-Za-z]", nh): - return + return None nh = "" key = _route_dedupe_key(rd=rd, net=net, nh=nh or "") if key in seen: - return + return None seen.add(key) metric, loc, tag, path = _split_rest(rest, path_continuation=path_continuation) - out.append( - { - "local_as": local_as[:16], - "afi": afi[:32], - "vrf": vrf[:128], - "neighbor": neighbor[:128], - "direction": direction[:8], - "rd": (rd or "")[:64], - "network": net[:128], - "next_hop": (nh or "")[:128], - "metric": metric[:32], - "loc_prf": loc[:32], - "tag": tag[:32], - "path": path[:256], - "status_codes": re.sub(r"\s+", "", (flags or "").strip())[:16], - "as_num": "", - "state": "", - "pfx_rcd": "", - } - ) + return { + "local_as": local_as[:16], + "afi": afi[:32], + "vrf": vrf[:128], + "neighbor": neighbor[:128], + "direction": direction[:8], + "rd": (rd or "")[:64], + "network": net[:128], + "next_hop": (nh or "")[:128], + "metric": metric[:32], + "loc_prf": loc[:32], + "tag": tag[:32], + "path": path[:256], + "status_codes": re.sub(r"\s+", "", (flags or "").strip())[:16], + "as_num": "", + "state": "", + "pfx_rcd": "", + } def _hand_parse( @@ -262,9 +259,8 @@ def _hand_parse( neighbor: str = "", direction: str = "", **_kw: Any, -) -> list[dict[str, Any]]: - """Parse neighbor in/out tables; join heavy IPv6 / From / metric wraps.""" - out: list[dict[str, Any]] = [] +): + """Yield neighbor in/out route rows (streaming; joins heavy IPv6 / From wraps).""" seen: set[str] = set() pending_net = "" pending_flags = "" @@ -273,13 +269,12 @@ def _hand_parse( current_rd = "" 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 if not pending_net: - return + return None merged = " ".join(x for x in (pending_rest, rest) if x).strip() - _emit_route( - out, + row = _emit_route( seen, local_as=local_as, afi=afi, @@ -297,20 +292,24 @@ def _hand_parse( pending_flags = "" pending_nh = "" 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 - _flush_pending() + flushed = _flush_pending() pending_net = net pending_flags = flags pending_nh = nh pending_rest = rest + return flushed for raw in _normalize_cli_text(raw_text).splitlines(): line = raw.rstrip() rd_m = _RD_RE.match(line.strip()) if rd_m: - _flush_pending() + row = _flush_pending() + if row: + yield row current_rd = (rd_m.group("rd") or "").strip() vrf_from_rd = (rd_m.group("vrf") or "").strip() if vrf_from_rd: @@ -331,15 +330,21 @@ def _hand_parse( pending_nh = first more = " ".join(parts[1:]) 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: pending_rest = " ".join(x for x in (pending_rest, more) if x) continue 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 else: - _flush_pending(rest=tok, path_continuation=True) + row = _flush_pending(rest=tok, path_continuation=True) + if row: + yield row continue # 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")): rest = (m.group("rest") or "").strip() if _rest_looks_complete(rest): - _flush_pending() - _emit_route( - out, + row = _flush_pending() + if row: + yield row + row = _emit_route( seen, local_as=local_as, afi=afi, @@ -370,19 +376,27 @@ def _hand_parse( rest=rest, flags=flags, ) + if row: + yield row else: - # Empty rest or From-only → wait for metric/path wrap - _start_pending(net=m.group("net"), nh=m.group("nh"), flags=flags, rest=rest) + row = _start_pending( + net=m.group("net"), nh=m.group("nh"), flags=flags, rest=rest + ) + if row: + yield row continue # Network alone → wait for next-hop wrap m_net = _NET_ONLY_RE.match(body) 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 - _flush_pending() - return out + row = _flush_pending() + if row: + yield row def normalize_bgp_route( @@ -391,8 +405,8 @@ def normalize_bgp_route( command: str = "", params: dict[str, str] | None = None, **_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) if _empty_if_total_zero(raw_text): return [] diff --git a/netx_api/biz_state/spool.py b/netx_api/biz_state/spool.py index b0d00e9..367645c 100644 --- a/netx_api/biz_state/spool.py +++ b/netx_api/biz_state/spool.py @@ -6,6 +6,7 @@ import json import logging import re import shutil +from collections.abc import Iterable, Iterator, Mapping from dataclasses import dataclass, field from pathlib import Path 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) -def write_records(batch_id: str, cmd_id: str, records: list[dict[str, Any]]) -> str: - """Write parsed records as JSONL; return relative path.""" +def write_records( + 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) + n = 0 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)) + for rec in records or (): + if not isinstance(rec, Mapping): + continue + fh.write(json.dumps(dict(rec), ensure_ascii=False, default=str, separators=(",", ":"))) fh.write("\n") + n += 1 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: @@ -114,13 +123,26 @@ def read_raw_text(rel_path: str, *, max_bytes: int = 0) -> str: 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: - return [] + return 0 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]] = [] + return 0 + 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: for line in fh: line = line.strip() @@ -131,8 +153,27 @@ def read_records(rel_path: str) -> list[dict[str, Any]]: except json.JSONDecodeError: continue if isinstance(rec, dict): - out.append(rec) - return out + yield rec + + +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 diff --git a/tests/test_biz_state_spool.py b/tests/test_biz_state_spool.py index 7a1a59a..a6a40b2 100644 --- a/tests/test_biz_state_spool.py +++ b/tests/test_biz_state_spool.py @@ -17,6 +17,7 @@ from netx_api.biz_state import spool as spool_mod from netx_api.biz_state.spool import ( SpooledCommand, clear_batch_spool, + iter_record_chunks, read_raw_text, read_records, write_raw_text, @@ -43,11 +44,23 @@ class BizStateSpoolIoTests(unittest.TestCase): 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"}]) + rrel, n = write_records(bid, cid, [{"ip": "1.1.1.1"}, {"ip": "2.2.2.2"}]) + self.assertEqual(n, 2) recs = read_records(rrel) self.assertEqual(len(recs), 2) 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: bid = "b2" cid = "c2" @@ -116,11 +129,12 @@ class BizStateFlushSpoolTests(unittest.TestCase): 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( + rec_rel, rec_n = write_records( "b-spool", cid, [{"ip": "10.0.0.1", "mac": "aaaa"}, {"ip": "10.0.0.2", "mac": "bbbb"}], ) + self.assertEqual(rec_n, 2) pending = [ SpooledCommand( id=cid, @@ -134,6 +148,7 @@ class BizStateFlushSpoolTests(unittest.TestCase): message="spooled", raw_rel_path=raw_rel, records_rel_path=rec_rel, + row_count=rec_n, persist_kind="metric", ) ] @@ -158,6 +173,36 @@ class BizStateFlushSpoolTests(unittest.TestCase): self.assertEqual(batch.command_count, 1) 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: pending: list[SpooledCommand] = [] for i in range(5): diff --git a/tests/test_zte_extended_parsers.py b/tests/test_zte_extended_parsers.py index f705455..88c14e1 100644 --- a/tests/test_zte_extended_parsers.py +++ b/tests/test_zte_extended_parsers.py @@ -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.zte import ( normalize_bgp_peer, - normalize_bgp_route, + normalize_bgp_route as _normalize_bgp_route_stream, normalize_ip_route, normalize_ipv6_route, 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 +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 = """ Name Default RD Protocols VRF ID CUST_A 100:1 ipv4 1