diff --git a/netx_api/biz_state/retention.py b/netx_api/biz_state/retention.py index d14c554..fde4107 100644 --- a/netx_api/biz_state/retention.py +++ b/netx_api/biz_state/retention.py @@ -6,6 +6,7 @@ from collections import defaultdict from datetime import datetime, timedelta from typing import Any +from sqlalchemy import or_, select from sqlalchemy.orm import Session from ..models import ( @@ -26,38 +27,54 @@ def _utcnow() -> datetime: return datetime.utcnow() -def protected_batch_map(db: Session, *, task_id: str = "") -> dict[str, list[str]]: - """batch_id → reason codes. Empty task_id = all tasks.""" +def protected_batch_map( + db: Session, *, task_id: str = "", batch_ids: list[str] | None = None, +) -> dict[str, list[str]]: + """Only read reference columns, scoped to the requested task/batches when given.""" out: dict[str, list[str]] = defaultdict(list) + wanted = set(batch_ids) if batch_ids is not None else None + if wanted == set(): + return {} + scope = select(BizStateBatch.id) + q = db.query(BizStateBatch.id, BizStateBatch.is_baseline, BizStateBatch.status) + if task_id: + scope = scope.where(BizStateBatch.task_id == task_id) + q = q.filter(BizStateBatch.task_id == task_id) + if wanted is not None: + scope = scope.where(BizStateBatch.id.in_(wanted)) + q = q.filter(BizStateBatch.id.in_(wanted)) + batch_rows = q.all() if task_id or wanted is not None else q.filter(or_( + BizStateBatch.is_baseline.is_(True), BizStateBatch.status.in_(("queued", "running")), + )).all() + if task_id: + wanted = {row.id for row in batch_rows} def add(bid: str, reason: str) -> None: b = str(bid or "").strip() - if not b: + if not b or (wanted is not None and b not in wanted): return if reason not in out[b]: out[b].append(reason) - q = db.query(BizStateBatch) - if task_id: - q = q.filter(BizStateBatch.task_id == task_id) - for b in q.filter(BizStateBatch.is_baseline.is_(True)).all(): - add(b.id, "manual_baseline") + for b in batch_rows: + if b.is_baseline: + add(b.id, "manual_baseline") + if b.status in ("queued", "running"): + add(b.id, "active_collection") - for j in db.query(BizCompareJob).all(): - add(j.before_batch_id, "compare_job_before") - add(j.after_batch_id, "compare_job_after") - - for r in db.query(BizCompareRun).all(): - add(r.before_batch_id, "compare_run_before") - add(r.after_batch_id, "compare_run_after") - - for p in db.query(BizMigrationProject).all(): - add(p.old_baseline_batch_id, "migration_old_baseline") - add(p.new_baseline_batch_id, "migration_new_baseline") - - for r in db.query(BizMigrationRun).all(): - add(r.old_batch_id, "migration_run_old") - add(r.new_batch_id, "migration_run_new") + references = ( + (BizCompareJob.before_batch_id, BizCompareJob.after_batch_id, "compare_job_before", "compare_job_after"), + (BizCompareRun.before_batch_id, BizCompareRun.after_batch_id, "compare_run_before", "compare_run_after"), + (BizMigrationProject.old_baseline_batch_id, BizMigrationProject.new_baseline_batch_id, "migration_old_baseline", "migration_new_baseline"), + (BizMigrationRun.old_batch_id, BizMigrationRun.new_batch_id, "migration_run_old", "migration_run_new"), + ) + for before, after, before_reason, after_reason in references: + refs = db.query(before, after) + if task_id or batch_ids is not None: + refs = refs.filter(or_(before.in_(scope), after.in_(scope))) + for before_id, after_id in refs: + add(before_id, before_reason) + add(after_id, after_reason) return dict(out) @@ -67,7 +84,7 @@ def protected_batch_ids(db: Session, *, task_id: str = "") -> set[str]: def batch_protect_info(db: Session, batch_id: str) -> dict[str, Any]: - reasons = protected_batch_map(db).get(batch_id, []) + reasons = protected_batch_map(db, batch_ids=[batch_id]).get(batch_id, []) b = db.get(BizStateBatch, batch_id) if b and bool(getattr(b, "is_baseline", False)) and "manual_baseline" not in reasons: reasons = ["manual_baseline", *reasons] diff --git a/netx_api/biz_state/service.py b/netx_api/biz_state/service.py index f3ad60a..da5797c 100644 --- a/netx_api/biz_state/service.py +++ b/netx_api/biz_state/service.py @@ -9,8 +9,8 @@ from typing import Any from uuid import uuid4 from fastapi import HTTPException -from sqlalchemy import String, cast, func, or_ -from sqlalchemy.orm import Session +from sqlalchemy import String, case, cast, func, or_ +from sqlalchemy.orm import Session, defer from ..lldp_shared import resolve_vendor_key from ..models import ( @@ -354,13 +354,15 @@ def get_task(db: Session, task_id: str) -> dict[str, Any]: .order_by(BizStateTaskItem.sort_order.asc()) .all() ) + bindings_by_item: dict[str, list[Any]] = {} + if items: + for binding in db.query(BizStateTaskItemBinding).filter( + BizStateTaskItemBinding.item_id.in_([it.id for it in items]) + ).all(): + bindings_by_item.setdefault(binding.item_id, []).append(binding) item_out = [] for it in items: - binds = ( - db.query(BizStateTaskItemBinding) - .filter(BizStateTaskItemBinding.item_id == it.id) - .all() - ) + binds = bindings_by_item.get(it.id, []) item_out.append( { "id": it.id, @@ -373,6 +375,11 @@ def get_task(db: Session, task_id: str) -> dict[str, Any]: "bindings": [{"placeholder": b.placeholder, "value": b.value} for b in binds], } ) + return {**_task_summary(task), "items": item_out} + + +def _task_summary(task: BizStateTask) -> dict[str, Any]: + """Task metadata without command items or bindings (safe for frequent polling).""" return { "id": task.id, "source": task.source, @@ -396,10 +403,16 @@ def get_task(db: Session, task_id: str) -> dict[str, Any]: if task.last_collect_ended_at else None, "last_error": task.last_error, - "items": item_out, } +def get_task_progress(db: Session, task_id: str) -> dict[str, Any]: + task = db.get(BizStateTask, task_id) + if not task: + raise HTTPException(status_code=404, detail="task_not_found") + return _task_summary(task) + + def list_tasks(db: Session, *, purpose: str | None = None) -> list[dict[str, Any]]: from sqlalchemy import or_ @@ -447,6 +460,8 @@ def delete_task(db: Session, task_id: str) -> None: task = db.get(BizStateTask, task_id) if not task: raise HTTPException(status_code=404, detail="task_not_found") + if task.collect_running: + raise HTTPException(status_code=409, detail="task_collecting") batches = db.query(BizStateBatch).filter(BizStateBatch.task_id == task_id).all() # Refuse if any batch is still referenced by compare/migration (manual baseline is OK to drop with task) pmap = protected_batch_map(db, task_id=task_id) @@ -500,7 +515,7 @@ def list_batches(db: Session, task_id: str, *, limit: int = 50) -> list[dict[str .limit(max(1, min(500, int(limit)))) .all() ) - pmap = protected_batch_map(db, task_id=task_id) + pmap = protected_batch_map(db, batch_ids=[b.id for b in rows]) out: list[dict[str, Any]] = [] for b in rows: reasons = list(pmap.get(b.id, [])) @@ -622,12 +637,29 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]: b = db.get(BizStateBatch, batch_id) if not b: raise HTTPException(status_code=404, detail="batch_not_found") - cmds = ( - db.query(BizStateBatchCommand) - .filter(BizStateBatchCommand.batch_id == batch_id) - .order_by(BizStateBatchCommand.created_at.asc()) - .all() + raw = func.coalesce(BizStateBatchCommand.raw_text, "") + raw_len = func.length(raw) + # Compute legacy counts in SQL; workbook summaries never transfer raw CLI text. + line_count = case( + (BizStateBatchCommand.raw_line_count > 0, BizStateBatchCommand.raw_line_count), + (raw_len == 0, 0), + else_=raw_len - func.length(func.replace(raw, "\n", "")) + + case((func.substr(raw, raw_len, 1) == "\n", 0), else_=1), ) + whitespace = " \t\r\n\v\f" + trimmed = (func.btrim(raw, whitespace) if db.get_bind().dialect.name == "postgresql" + else func.trim(raw, whitespace)) + cmd_raw_stats: dict[str, tuple[int, bool]] = {} + cmds = [] + for c, lines, has_raw in ( + db.query(BizStateBatchCommand, line_count, func.length(trimmed) > 0) + .options(defer(BizStateBatchCommand.raw_text)) + .filter(BizStateBatchCommand.batch_id == batch_id) + .order_by(BizStateBatchCommand.created_at.asc(), BizStateBatchCommand.id.asc()) + .all() + ): + cmds.append(c) + cmd_raw_stats[c.id] = (int(lines or 0), bool(has_raw)) protect = batch_protect_info(db, batch_id) # Per-metric row counts (generic table) @@ -668,11 +700,14 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]: # Prefer primary collect rows over aux / aux_cached for the same CLI sheet_cmds[id_].append(cmd_info) - primary_cmds_by_cli: dict[str, dict[str, Any]] = {} + # Aux failures must remain visible even if their primary command was stored later. + primary_cmds_by_cli = { + normalize_command(str(c.raw_command or "")) for c in cmds + if not str(c.parse_status or "").strip().lower().startswith("aux") + } for c in cmds: status = str(c.parse_status or "").strip().lower() is_aux = status.startswith("aux") - raw = c.raw_text or "" info = { "id": c.id, "profile_id": c.profile_id, @@ -682,15 +717,13 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]: "params": c.params_json or {}, "parse_status": c.parse_status, "row_count": c.row_count, - "raw_line_count": _cmd_raw_line_count(c), + "raw_line_count": cmd_raw_stats[c.id][0], "declared_total": _cmd_declared_total(c), "message": c.message, - "has_raw": bool(str(raw).strip()), + "has_raw": cmd_raw_stats[c.id][1], "is_aux": is_aux, } cmd_n = normalize_command(str(c.raw_command or "")) - if not is_aux and cmd_n: - primary_cmds_by_cli.setdefault(cmd_n, info) # Commands sheet: hide successful aux when the same CLI already has a primary row; # keep failed/skipped aux visible so partial reasons are not hidden. if is_aux and cmd_n and cmd_n in primary_cmds_by_cli: @@ -746,10 +779,10 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]: "params": c.params_json or {}, "parse_status": c.parse_status, "row_count": c.row_count, - "raw_line_count": _cmd_raw_line_count(c), + "raw_line_count": cmd_raw_stats[c.id][0], "declared_total": _cmd_declared_total(c), "message": c.message, - "has_raw": bool(str(c.raw_text or "").strip()), + "has_raw": cmd_raw_stats[c.id][1], "is_aux": True, } ) @@ -844,6 +877,7 @@ def list_batch_metric_rows( size_n = max(1, min(200, int(page_size or 50))) kw_n = str(kw or "").strip() col_n = str(column or "").strip() + like = "%" + kw_n.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%" fields = metric_field_map().get(mid) or [] columns = [ @@ -859,30 +893,30 @@ def list_batch_metric_rows( if mid == "lldp_neighbor": q = db.query(BizStateLldpNeighbor).filter(BizStateLldpNeighbor.batch_id == batch_id) if kw_n: - like = f"%{kw_n}%" if col_n == "local_if": - q = q.filter(BizStateLldpNeighbor.local_if.ilike(like)) + q = q.filter(BizStateLldpNeighbor.local_if.ilike(like, escape="\\")) elif col_n == "remote_sys": - q = q.filter(BizStateLldpNeighbor.remote_sys.ilike(like)) + q = q.filter(BizStateLldpNeighbor.remote_sys.ilike(like, escape="\\")) elif col_n == "remote_if": - q = q.filter(BizStateLldpNeighbor.remote_if.ilike(like)) + q = q.filter(BizStateLldpNeighbor.remote_if.ilike(like, escape="\\")) elif col_n == "remote_ip": - q = q.filter(BizStateLldpNeighbor.remote_ip.ilike(like)) + q = q.filter(BizStateLldpNeighbor.remote_ip.ilike(like, escape="\\")) elif col_n == "protocol": - q = q.filter(BizStateLldpNeighbor.protocol.ilike(like)) + q = q.filter(BizStateLldpNeighbor.protocol.ilike(like, escape="\\")) else: q = q.filter( or_( - BizStateLldpNeighbor.local_if.ilike(like), - BizStateLldpNeighbor.remote_sys.ilike(like), - BizStateLldpNeighbor.remote_if.ilike(like), - BizStateLldpNeighbor.remote_ip.ilike(like), - BizStateLldpNeighbor.protocol.ilike(like), + BizStateLldpNeighbor.local_if.ilike(like, escape="\\"), + BizStateLldpNeighbor.remote_sys.ilike(like, escape="\\"), + BizStateLldpNeighbor.remote_if.ilike(like, escape="\\"), + BizStateLldpNeighbor.remote_ip.ilike(like, escape="\\"), + BizStateLldpNeighbor.protocol.ilike(like, escape="\\"), ) ) total = int(q.count() or 0) rows_db = ( - q.order_by(BizStateLldpNeighbor.local_if.asc()) + q.order_by(BizStateLldpNeighbor.local_if.asc(), BizStateLldpNeighbor.remote_sys.asc(), + BizStateLldpNeighbor.remote_if.asc(), BizStateLldpNeighbor.id.asc()) .offset((page_n - 1) * size_n) .limit(size_n) .all() @@ -911,12 +945,11 @@ def list_batch_metric_rows( BizStateMetricRow.metric_id == mid, ) if kw_n: - like = f"%{kw_n}%" if col_n: # JSON path as text — works on Postgres JSONB and SQLite JSON - q = q.filter(cast(BizStateMetricRow.data_json[col_n], String).ilike(like)) + q = q.filter(cast(BizStateMetricRow.data_json[col_n].as_string(), String).ilike(like, escape="\\")) else: - q = q.filter(cast(BizStateMetricRow.data_json, String).ilike(like)) + q = q.filter(cast(BizStateMetricRow.data_json, String).ilike(like, escape="\\")) total = int(q.count() or 0) rows_db = ( q.order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc()) @@ -1019,56 +1052,64 @@ def export_batch_zip(db: Session, batch_id: str) -> bytes: db.query(BizStateLldpNeighbor) .filter(BizStateLldpNeighbor.batch_id == batch_id) .order_by(BizStateLldpNeighbor.local_if.asc()) - .all() + .yield_per(1000) ) - csv_lines = ["local_if,remote_sys,remote_if,remote_ip,protocol"] - for n in neighbors: - csv_lines.append( - ",".join( - [ - _csv(n.local_if), - _csv(n.remote_sys), - _csv(n.remote_if), - _csv(n.remote_ip), - _csv(n.protocol), - ] - ) - ) - zf.writestr("tables/lldp_neighbor.csv", "\n".join(csv_lines) + "\n") + with zf.open("tables/lldp_neighbor.csv", "w") as dest: + dest.write(b"local_if,remote_sys,remote_if,remote_ip,protocol\n") + for n in neighbors: + line = ",".join(_csv(v) for v in ( + n.local_if, n.remote_sys, n.remote_if, n.remote_ip, n.protocol, + )) + "\n" + dest.write(line.encode("utf-8")) # Generic metrics CSV (stream by metric_id) for sheet in detail.get("sheets") or []: mid = str(sheet.get("metric_id") or "").strip() if not mid or mid == "lldp_neighbor": continue - rows = ( - db.query(BizStateMetricRow) - .filter( - BizStateMetricRow.batch_id == batch_id, - BizStateMetricRow.metric_id == mid, - ) - .order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc()) - .all() - ) - if not rows: - continue - recs = [dict(r.data_json or {}) for r in rows] - cols: list[str] = [] - for rec in recs: + # First scan discovers the same complete, ordered header as the legacy export. + # Second scan writes rows without materializing the table or CSV. + col_keys: dict[str, None] = {} + has_rows = False + for rec in _iter_metric_export_rows(db, batch_id, mid): + has_rows = True for k in rec.keys(): - if k not in cols: - cols.append(str(k)) - out_lines = [",".join(_csv(c) for c in cols)] - for rec in recs: - out_lines.append(",".join(_csv(str(rec.get(c, "") or "")) for c in cols)) + col_keys.setdefault(str(k), None) + if not has_rows: + continue + cols = list(col_keys) safe = "".join(ch if ch.isalnum() or ch in "-_" else "_" for ch in mid)[:80] or "metric" - zf.writestr(f"tables/{safe}.csv", "\n".join(out_lines) + "\n") + with zf.open(f"tables/{safe}.csv", "w") as dest: + dest.write((",".join(_csv(c) for c in cols) + "\n").encode("utf-8")) + for rec in _iter_metric_export_rows(db, batch_id, mid): + line = ",".join(_csv(rec.get(c)) for c in cols) + "\n" + dest.write(line.encode("utf-8")) return buf.getvalue() -def _csv(v: str) -> str: - s = str(v or "") - if any(ch in s for ch in ",\"\n"): +def _iter_metric_export_rows(db: Session, batch_id: str, metric_id: str, chunk_size: int = 1000): + cursor: tuple[int, str] | None = None + while True: + query = db.query(BizStateMetricRow.seq, BizStateMetricRow.id, BizStateMetricRow.data_json).filter( + BizStateMetricRow.batch_id == batch_id, BizStateMetricRow.metric_id == metric_id, + ) + if cursor is not None: + seq, row_id = cursor + query = query.filter(or_( + BizStateMetricRow.seq > seq, + (BizStateMetricRow.seq == seq) & (BizStateMetricRow.id > row_id), + )) + chunk = query.order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc()).limit(chunk_size).all() + if not chunk: + return + for row in chunk: + yield row.data_json or {} + cursor = (chunk[-1].seq, chunk[-1].id) + + +def _csv(v: Any) -> str: + s = "" if v is None else str(v) + if any(ch in s for ch in ",\"\r\n"): return '"' + s.replace('"', '""') + '"' return s diff --git a/netx_api/biz_state_router.py b/netx_api/biz_state_router.py index 20a0319..9c59382 100644 --- a/netx_api/biz_state_router.py +++ b/netx_api/biz_state_router.py @@ -173,6 +173,11 @@ def api_get_task(task_id: str, db: Session = Depends(get_db)) -> dict[str, Any]: return svc.get_task(db, task_id) +@router.get("/tasks/{task_id}/progress") +def api_task_progress(task_id: str, db: Session = Depends(get_db)) -> dict[str, Any]: + return svc.get_task_progress(db, task_id) + + @router.get("/tasks/{task_id}/commands") def api_plan_task_commands( task_id: str, diff --git a/tests/test_biz_state_monitor_service.py b/tests/test_biz_state_monitor_service.py new file mode 100644 index 0000000..6f266cd --- /dev/null +++ b/tests/test_biz_state_monitor_service.py @@ -0,0 +1,144 @@ +"""Monitoring query correctness, bounded reads and active-collection protection.""" +import csv +import io +import zipfile +from datetime import datetime, timedelta + +import pytest +from fastapi import HTTPException +from sqlalchemy import create_engine, event +from sqlalchemy.orm import sessionmaker + +from netx_api.db import Base +from netx_api.models import ( + BizCompareRun, BizStateBatch, BizStateBatchCommand, BizStateLldpNeighbor, + BizStateMetricRow, BizStateTask, BizStateTaskItem, BizStateTaskItemBinding, +) +from netx_api.biz_state import service +from netx_api.biz_state.retention import protected_batch_map, purge_task_batches + + +@pytest.fixture +def db(): + engine = create_engine("sqlite+pysqlite:///:memory:") + Base.metadata.create_all(engine) + session = sessionmaker(bind=engine, expire_on_commit=False)() + session.add(BizStateTask(id="t", ne_id="ne", retention_days=1)) + session.add(BizStateBatch(id="b", task_id="t", status="success")) + session.commit() + yield session + session.close() + engine.dispose() + + +def test_task_detail_uses_one_binding_query_and_progress_omits_items(db): + for n in range(30): + db.add(BizStateTaskItem(id=f"i{n}", task_id="t", sort_order=n)) + for ph, val in (("vrf", f"v{n}"), ("neighbor", f"p{n}")): + db.add(BizStateTaskItemBinding(id=f"{n}-{ph}", item_id=f"i{n}", placeholder=ph, value=val)) + db.commit() + db.expunge_all() + statements = [] + event.listen(db.get_bind(), "before_cursor_execute", lambda c, cur, sql, p, ctx, many: statements.append(sql)) + result = service.get_task(db, "t") + assert len(statements) == 3 + assert len(result["items"]) == 30 + assert result["items"][9]["sort_order"] == 9 + assert {b["value"] for b in result["items"][9]["bindings"]} == {"v9", "p9"} + db.expunge_all() + statements.clear() + progress = service.get_task_progress(db, "t") + assert len(statements) == 1 + assert "items" not in progress + assert progress["id"] == "t" + with pytest.raises(HTTPException) as exc: + service.get_task_progress(db, "missing") + assert exc.value.status_code == 404 + + +@pytest.mark.parametrize("metric,column", [("custom", "name"), ("lldp_neighbor", "remote_sys")]) +@pytest.mark.parametrize("keyword", ["%", "_", "\\"]) +def test_metric_search_treats_wildcards_literally(db, metric, column, keyword): + for n, value in enumerate((f"exact{keyword}value", "different-value")): + if metric == "lldp_neighbor": + db.add(BizStateLldpNeighbor(id=f"r{n}", batch_id="b", local_if=f"ge{n}", remote_sys=value)) + else: + db.add(BizStateMetricRow(id=f"r{n}", batch_id="b", metric_id=metric, seq=n, data_json={column: value})) + db.commit() + result = service.list_batch_metric_rows(db, "b", metric, kw=keyword, column=column) + assert result["total"] == 1 + assert result["items"][0][column] == f"exact{keyword}value" + + +def test_summary_does_not_select_raw_text_and_keeps_early_aux_failure(db): + now = datetime(2026, 10, 10) + for cid, status, raw, stored in ( + ("a", "aux_failed", "first\nsecond\n", 0), + ("b", "ok", "X" * 1_000_000, 999), + ("c", "aux_cached", " \t\n", 0), + ): + db.add(BizStateBatchCommand(id=cid, batch_id="b", raw_command="show arp" if cid != "c" else "show config", + metric_id="arp", parse_status=status, raw_text=raw, raw_line_count=stored, created_at=now)) + db.commit() + db.expunge_all() + column_names = [] + event.listen(db.get_bind(), "after_cursor_execute", lambda c, cur, sql, p, ctx, many: + column_names.extend([d[0] for d in (cur.description or [])]) if "biz_state_batch_command" in sql else None) + summary = service.get_batch(db, "b") + cmds = {c["id"]: c for c in summary["commands"]} + assert "biz_state_batch_command_raw_text" not in column_names + assert cmds["a"]["raw_line_count"] == 2 + assert cmds["b"]["raw_line_count"] == 999 + assert cmds["b"]["has_raw"] is True + assert cmds["c"]["has_raw"] is False + assert summary["command_stats"]["aux_failed"] == 1 + assert service.get_batch_command(db, "b", "b")["raw_text"] == "X" * 1_000_000 + + +def test_active_batches_cannot_be_deleted_or_purged(db): + for status in ("queued", "running"): + db.add(BizStateBatch(id=status, task_id="t", status=status, started_at=datetime.utcnow() - timedelta(days=3))) + db.commit() + for bid in ("queued", "running"): + with pytest.raises(HTTPException) as exc: + service.delete_batch(db, bid) + assert exc.value.status_code == 409 + assert "active_collection" in exc.value.detail["reasons"] + result = service.delete_batches_bulk(db, ["queued", "running"]) + assert result["deleted"] == [] + purge_task_batches(db, db.get(BizStateTask, "t")) + assert db.get(BizStateBatch, "queued") is not None + assert db.get(BizStateBatch, "running") is not None + db.get(BizStateTask, "t").collect_running = True + db.commit() + with pytest.raises(HTTPException) as exc: + service.delete_task(db, "t") + assert exc.value.detail == "task_collecting" + + +def test_protection_reads_only_scoped_reference_columns(db): + db.add(BizStateBatch(id="other", task_id="other-task", status="success")) + db.add(BizCompareRun(id="r", before_batch_id="b", after_batch_id="other", summary_json={"large": "x" * 10000})) + db.commit() + db.expunge_all() + statements = [] + event.listen(db.get_bind(), "before_cursor_execute", lambda c, cur, sql, p, ctx, many: statements.append(sql)) + assert protected_batch_map(db, batch_ids=["b"]) == {"b": ["compare_run_before"]} + assert all("summary_json" not in sql for sql in statements) + assert protected_batch_map(db, task_id="t") == {"b": ["compare_run_before"]} + assert protected_batch_map(db, batch_ids=[]) == {} + + +def test_export_preserves_false_zero_csv_and_duplicate_sequences(db): + for n in range(1005): + rec = {"n": n, "enabled": False, "note": 'a,b"c\rd\ne' if n == 0 else None} + if n == 1004: + rec["late"] = "last-column" + db.add(BizStateMetricRow(id=f"r{n:04}", batch_id="b", metric_id="custom", seq=0, data_json=rec)) + db.commit() + with zipfile.ZipFile(io.BytesIO(service.export_batch_zip(db, "b"))) as archive: + rows = list(csv.DictReader(io.StringIO(archive.read("tables/custom.csv").decode("utf-8"), newline=""))) + assert len(rows) == 1005 + assert rows[0] == {"n": "0", "enabled": "False", "note": 'a,b"c\rd\ne', "late": ""} + assert rows[-1]["late"] == "last-column" + assert [row["n"] for row in rows] == [str(n) for n in range(1005)]