From ba16167c8e949dc8bb1215cf9a829d1140fc8e9b Mon Sep 17 00:00:00 2001 From: oliver Date: Thu, 8 Oct 2026 15:16:40 +0800 Subject: [PATCH] Add SQL JOIN compare for pushdown-safe sheets and mid-run result viewing. Co-authored-by: Cursor --- netx_api/biz_state/compare_service.py | 187 ++++++- netx_api/biz_state/compare_sql.py | 615 +++++++++++++++++++++++ tests/test_biz_state_compare_sql.py | 215 ++++++++ web/src/pages/network/BizComparePage.tsx | 17 +- 4 files changed, 1004 insertions(+), 30 deletions(-) create mode 100644 netx_api/biz_state/compare_sql.py create mode 100644 tests/test_biz_state_compare_sql.py diff --git a/netx_api/biz_state/compare_service.py b/netx_api/biz_state/compare_service.py index 2507d06..2f8b2f5 100644 --- a/netx_api/biz_state/compare_service.py +++ b/netx_api/biz_state/compare_service.py @@ -453,6 +453,53 @@ def _sheet_meta_from_summary(summary: dict[str, Any], run: BizCompareRun, tpl: A ] +def _metric_row_estimate( + db: Session, batch_ids: list[str], metric_id: str +) -> int: + """Cheap size hint from batch command row_count (max across sides).""" + mid = str(metric_id or "").strip() + if not mid: + return 0 + best = 0 + for bid in batch_ids: + bid = str(bid or "").strip() + if not bid: + continue + rows = ( + db.query(BizStateBatchCommand) + .filter( + BizStateBatchCommand.batch_id == bid, + BizStateBatchCommand.metric_id == mid, + ) + .all() + ) + if not rows: + continue + n = sum(int(c.row_count or 0) for c in rows) + if n > best: + best = n + return best + + +def _order_sheets_small_first( + db: Session, + sheets: list[dict[str, Any]], + *, + before_batch_id: str, + after_batch_id: str, +) -> list[dict[str, Any]]: + """Run smaller metrics first so field engineers can review early results.""" + if len(sheets) <= 1: + return list(sheets) + batches = [before_batch_id, after_batch_id] + scored: list[tuple[int, int, dict[str, Any]]] = [] + for i, sheet in enumerate(sheets): + n = _metric_row_estimate(db, batches, str(sheet.get("metric_id") or "")) + scored.append((n, i, sheet)) + scored.sort(key=lambda x: (x[0], x[1])) + return [s for _, _, s in scored] + + def _pending_sheet_meta(sheet: dict[str, Any]) -> dict[str, Any]: """Placeholder meta so the UI lists all check items while a run is in flight.""" key_fields = list(sheet.get("key_fields") or []) @@ -1672,6 +1719,40 @@ def _resolve_after_batch(db: Session, job: BizCompareJob) -> str: return str(latest.id) if latest else "" +def _sheet_result_envelope( + sheet: dict[str, Any], + *, + key_fields: list[str], + iface_fields: list[str], + compare_fields: list[str], + display_fields: list[str], + row_filters: list[Any], + field_rules: list[Any], + ignore_ports: bool | None, + mode: str, + summary: dict[str, Any], + diffs: list[Any], + mapping_stats: dict[str, Any], +) -> dict[str, Any]: + return { + "sheet_id": sheet_key(sheet), + "title": sheet_title(sheet), + "metric_id": sheet["metric_id"], + "key_fields": key_fields, + "iface_fields": iface_fields, + "compare_fields": compare_fields, + "display_fields": display_fields, + "row_filters": row_filters, + "field_rules": field_rules, + "ignore_port_changes": ignore_ports, + "mode": mode, + "status": "done", + "summary": summary, + "diffs": diffs, + "mapping_stats": mapping_stats, + } + + def _run_sheet( db: Session, *, @@ -1704,6 +1785,58 @@ def _run_sheet( ignore_ports = bool(ignore_ports) mid = sheet["metric_id"] + # PostgreSQL path: pushdown-safe sheets join in-DB (BGP-scale). + from .compare_sql import can_sql_compare, run_sql_sheet_compare + + if can_sql_compare( + db, + sheet, + port_map=port_map, + iface_normalize_rules=iface_normalize_rules, + ): + try: + + def _sql_progress(side: str, n: int) -> None: + if on_load_progress: + on_load_progress(side, n) + + result = run_sql_sheet_compare( + db, + sheet=sheet, + before_batch_id=before_batch_id, + after_batch_id=after_batch_id, + store_unchanged=store_unchanged, + on_progress=_sql_progress, + ) + summary = dict(result["summary"]) + # Engine already sets raw counts / policy; keep keys stable + if "unchanged_policy" not in summary: + summary["unchanged_policy"] = resolve_unchanged_policy( + store_unchanged, + before_n=int(summary.get("before_count") or 0), + after_n=int(summary.get("after_count") or 0), + ) + return _sheet_result_envelope( + sheet, + key_fields=key_fields, + iface_fields=iface_fields, + compare_fields=compare_fields, + display_fields=display_fields, + row_filters=row_filters, + field_rules=field_rules, + ignore_ports=ignore_ports, + mode=mode, + summary=summary, + diffs=list(result.get("diffs") or []), + mapping_stats=dict(result.get("mapping_stats") or {}), + ) + except Exception: + _log.exception( + "sql compare fallback sheet=%s metric=%s — using Python engine", + sheet_key(sheet), + mid, + ) + def _before_chunk(n: int) -> None: if on_load_progress: on_load_progress("before", n) @@ -1742,23 +1875,21 @@ def _run_sheet( summary["after_raw_count"] = len(after_raw) summary["row_filters"] = len(row_filters) summary["unchanged_policy"] = policy - return { - "sheet_id": sheet_key(sheet), - "title": sheet_title(sheet), - "metric_id": sheet["metric_id"], - "key_fields": key_fields, - "iface_fields": iface_fields, - "compare_fields": compare_fields, - "display_fields": display_fields, - "row_filters": row_filters, - "field_rules": field_rules, - "ignore_port_changes": ignore_ports, - "mode": mode, - "status": "done", - "summary": summary, - "diffs": result["diffs"], - "mapping_stats": result["mapping_stats"], - } + summary.setdefault("engine", "python") + return _sheet_result_envelope( + sheet, + key_fields=key_fields, + iface_fields=iface_fields, + compare_fields=compare_fields, + display_fields=display_fields, + row_filters=row_filters, + field_rules=field_rules, + ignore_ports=ignore_ports, + mode=mode, + summary=summary, + diffs=list(result.get("diffs") or []), + mapping_stats=dict(result.get("mapping_stats") or {}), + ) def _validate_compare_job( @@ -1784,6 +1915,12 @@ def _validate_compare_job( sheets_cfg = _filter_enabled_sheets(sheets_cfg, getattr(j, "enabled_sheet_ids", None)) if not sheets_cfg: raise HTTPException(status_code=400, detail="no_enabled_sheets") + sheets_cfg = _order_sheets_small_first( + db, + sheets_cfg, + before_batch_id=before_batch_id, + after_batch_id=after_batch_id, + ) return { "job": j, "template": tpl, @@ -1910,6 +2047,12 @@ def _execute_compare_into_run(db: Session, run_id: str) -> dict[str, Any]: run.message = "no_enabled_sheets" db.commit() raise HTTPException(status_code=400, detail="no_enabled_sheets") + sheets_cfg = _order_sheets_small_first( + db, + sheets_cfg, + before_batch_id=before_batch_id, + after_batch_id=after_batch_id, + ) started_mono = time.monotonic() store_mode = normalize_store_unchanged(getattr(j, "store_unchanged", None)) @@ -1931,10 +2074,12 @@ def _execute_compare_into_run(db: Session, run_id: str) -> dict[str, Any]: field_counts: dict[str, int] = {} total = len(sheets_cfg) - # Prefer seeded pending sheets from create; rebuild if missing - sheet_metas: list[dict[str, Any]] = list((run.summary_json or {}).get("sheets") or []) - if len(sheet_metas) != total: - sheet_metas = [_pending_sheet_meta(s) for s in sheets_cfg] + # Prefer seeded pending sheets from create; realign to small-first order + prev_metas = list((run.summary_json or {}).get("sheets") or []) + by_key = {sheet_key(m): m for m in prev_metas} + sheet_metas = [ + by_key.get(sheet_key(s)) or _pending_sheet_meta(s) for s in sheets_cfg + ] def _publish_sheets() -> None: prev = dict(run.summary_json or {}) diff --git a/netx_api/biz_state/compare_sql.py b/netx_api/biz_state/compare_sql.py new file mode 100644 index 0000000..e66475b --- /dev/null +++ b/netx_api/biz_state/compare_sql.py @@ -0,0 +1,615 @@ +"""PostgreSQL-backed sheet compare (FULL OUTER JOIN) for pushdown-safe sheets. + +Falls back to the Python engine when port-map / complex rules cannot be expressed +in SQL. See ``can_sql_compare``. +""" + +from __future__ import annotations + +import logging +import re +from typing import Any, Callable, Mapping, Sequence +from uuid import uuid4 + +from sqlalchemy import text +from sqlalchemy.orm import Session + +from .compare_rules import effective_compare_fields, field_rule_map + +_log = logging.getLogger("netx.biz_state.compare_sql") + +_FIELD_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") +_SQL_FILTER_OPS = frozenset( + {"eq", "==", "ne", "!=", "in", "not_in", "nin", "contains", "empty", "not_empty", "nonempty", "ci_eq"} +) +_SQL_NORMALIZE = frozenset({"", "none", "strip", "lower", "upper", "empty_as_blank"}) +_SQL_COMPARE_MODES = frozenset({"", "eq", "ignore", "skip", "off"}) +_EMPTY_AS_BLANK = ("n/a", "na", "-", "--", "none", "null") + + +def _dialect_is_postgres(db: Session) -> bool: + bind = db.get_bind() + if bind is None: + return False + return str(getattr(bind.dialect, "name", "") or "").lower() in ("postgresql", "postgres") + + +def _safe_field(name: str) -> str: + f = str(name or "").strip() + if not f or not _FIELD_RE.match(f): + raise ValueError(f"unsafe_json_field:{name!r}") + return f + + +def _json_text_expr(alias: str, field: str) -> str: + """SQL expression: trimmed text from JSONB column ``alias.data`` / ``data_json``.""" + f = _safe_field(field) + # alias is caller-controlled (_b / data_json) — not user input + return f"trim(both from coalesce({alias}->>'{f}', ''))" + + +def _norm_expr(alias: str, field: str, normalize: str) -> str: + base = _json_text_expr(alias, field) + mode = str(normalize or "strip").strip().lower() or "strip" + if mode in ("", "none", "strip"): + return base + if mode == "lower": + return f"lower({base})" + if mode == "upper": + return f"upper({base})" + if mode == "empty_as_blank": + opts = ", ".join(f"'{x}'" for x in _EMPTY_AS_BLANK) + return ( + f"CASE WHEN lower({base}) IN ({opts}) THEN '' ELSE {base} END" + ) + raise ValueError(f"unsupported_normalize:{mode}") + + +def _filters_sql_compatible(filters: Sequence[Mapping[str, Any]] | None) -> bool: + fl = [f for f in (filters or []) if isinstance(f, dict)] + return all(_filter_node_compatible(f) for f in fl) + + +def _filter_node_compatible(filt: Mapping[str, Any]) -> bool: + if "any" in filt: + kids = filt.get("any") or [] + return isinstance(kids, list) and all( + isinstance(k, dict) and _filter_node_compatible(k) for k in kids + ) + if "all" in filt: + kids = filt.get("all") or [] + return isinstance(kids, list) and all( + isinstance(k, dict) and _filter_node_compatible(k) for k in kids + ) + field = str(filt.get("field") or "").strip() + op = str(filt.get("op") or "eq").strip().lower() + if op not in _SQL_FILTER_OPS: + return False + if op in ("empty", "not_empty", "nonempty"): + return bool(field) and bool(_FIELD_RE.match(field)) + if not field or not _FIELD_RE.match(field): + return False + if op in ("in", "not_in", "nin"): + vals = filt.get("value") + if vals is None: + return True + if not isinstance(vals, (list, tuple)): + vals = [vals] + return all(isinstance(x, (str, int, float, bool)) or x is None for x in vals) + return True + + +def _field_rules_sql_compatible(rules: Sequence[Mapping[str, Any]] | None) -> bool: + for raw in rules or []: + if not isinstance(raw, dict): + return False + name = str(raw.get("field") or "").strip() + if name and not _FIELD_RE.match(name): + return False + mode = str(raw.get("compare") or "eq").strip().lower() or "eq" + if mode not in _SQL_COMPARE_MODES: + return False + if raw.get("ignore") is True: + continue + if mode in ("ignore", "skip", "off"): + continue + norm = str(raw.get("normalize") or "strip").strip().lower() or "strip" + if norm not in _SQL_NORMALIZE: + return False + return True + + +def can_sql_compare( + db: Session, + sheet: Mapping[str, Any], + *, + port_map: Mapping[str, str] | None = None, + iface_normalize_rules: Sequence[Mapping[str, str]] | None = None, +) -> bool: + """True when this sheet can run entirely as a PostgreSQL JOIN.""" + if not _dialect_is_postgres(db): + return False + if port_map: + return False + key_fields = [str(k).strip() for k in (sheet.get("key_fields") or []) if str(k).strip()] + if not key_fields or any(not _FIELD_RE.match(k) for k in key_fields): + return False + iface_fields = [str(f).strip() for f in (sheet.get("iface_fields") or []) if str(f).strip()] + iface_set = set(iface_fields) + # Port-rename heuristic / map rewrite cannot be expressed here + ignore_ports = sheet.get("ignore_port_changes") + if ignore_ports is True: + return False + if ignore_ports is None and iface_set and any(k in iface_set for k in key_fields): + # Auto path may drop iface from match key — stay on Python + return False + field_rules = list(sheet.get("field_rules") or []) + if not _field_rules_sql_compatible(field_rules): + return False + compare_fields = effective_compare_fields( + list(sheet.get("compare_fields") or []), + field_rules, + ) + if any(not _FIELD_RE.match(str(f).strip()) for f in compare_fields if str(f).strip()): + return False + used = set(key_fields) | {str(f).strip() for f in compare_fields if str(f).strip()} + if iface_normalize_rules and (used & iface_set): + return False + if not _filters_sql_compatible(list(sheet.get("row_filters") or [])): + return False + if not str(sheet.get("metric_id") or "").strip(): + return False + return True + + +def compile_row_filters_sql( + filters: Sequence[Mapping[str, Any]] | None, + *, + json_col: str = "data_json", + param_prefix: str = "f", +) -> tuple[str, dict[str, Any]]: + """Compile template row_filters to SQL AND-clause + bind params. + + Returns ``(sql_fragment, params)``. Empty filters → ``(\"TRUE\", {})``. + ``json_col`` is the JSONB column/expression name (trusted). + """ + fl = [f for f in (filters or []) if isinstance(f, dict)] + if not fl: + return "TRUE", {} + params: dict[str, Any] = {} + counter = {"n": 0} + + def _next(name: str) -> str: + counter["n"] += 1 + return f"{param_prefix}_{counter['n']}_{name}" + + def _node(filt: Mapping[str, Any]) -> str: + if "any" in filt: + kids = [k for k in (filt.get("any") or []) if isinstance(k, dict)] + if not kids: + return "TRUE" + return "(" + " OR ".join(_node(k) for k in kids) + ")" + if "all" in filt: + kids = [k for k in (filt.get("all") or []) if isinstance(k, dict)] + if not kids: + return "TRUE" + return "(" + " AND ".join(_node(k) for k in kids) + ")" + field = _safe_field(str(filt.get("field") or "")) + op = str(filt.get("op") or "eq").strip().lower() + expr = _json_text_expr(json_col, field) + # _json_text_expr uses alias->> ; for bare column use data_json directly + if json_col == "data_json": + expr = f"trim(both from coalesce(data_json->>'{field}', ''))" + if op in ("empty",): + return f"({expr} = '')" + if op in ("not_empty", "nonempty"): + return f"({expr} <> '')" + expect = filt.get("value") + if op in ("eq", "==", "ci_eq"): + key = _next("v") + params[key] = str(expect or "").strip() + if op == "ci_eq" or op in ("eq", "=="): + # Python eq is case-insensitive via .lower() + params[key] = str(expect or "").strip().lower() + return f"(lower({expr}) = :{key})" + if op in ("ne", "!="): + key = _next("v") + params[key] = str(expect or "").strip().lower() + return f"(lower({expr}) <> :{key})" + if op == "contains": + key = _next("v") + needle = str(expect or "").strip().lower() + params[key] = f"%{needle}%" + return f"(lower({expr}) LIKE :{key})" + if op in ("in", "not_in", "nin"): + vals = expect + if vals is None: + vals = [] + if not isinstance(vals, (list, tuple)): + vals = [vals] + clean = [str(x).strip().lower() for x in vals if str(x).strip()] + if not clean: + return "FALSE" if op == "in" else "TRUE" + keys = [] + for i, v in enumerate(clean): + k = _next(f"in{i}") + params[k] = v + keys.append(f":{k}") + inside = f"lower({expr}) IN ({', '.join(keys)})" + return f"({inside})" if op == "in" else f"(NOT {inside})" + raise ValueError(f"unsupported_filter_op:{op}") + + parts = [_node(f) for f in fl] + return "(" + " AND ".join(parts) + ")", params + + +def _rk_sql(key_fields: list[str], *, json_col: str = "data_json") -> str: + parts = [] + for k in key_fields: + f = _safe_field(k) + if json_col == "data_json": + parts.append(f"trim(both from coalesce(data_json->>'{f}', ''))") + else: + parts.append(f"trim(both from coalesce({json_col}->>'{f}', ''))") + if len(parts) == 1: + return parts[0] + return "concat_ws('|', " + ", ".join(parts) + ")" + + +def _changed_predicate( + compare_fields: list[str], + rules: dict[str, dict[str, Any]], + *, + before_alias: str = "b", + after_alias: str = "a", +) -> str: + """SQL boolean: True when any compare field differs (IS DISTINCT FROM).""" + if not compare_fields: + return "FALSE" + clauses: list[str] = [] + for f in compare_fields: + name = str(f).strip() + if not name: + continue + rule = rules.get(name) or {} + mode = str(rule.get("compare") or "eq").strip().lower() or "eq" + if mode in ("ignore", "skip", "off") or rule.get("ignore") is True: + continue + norm = str(rule.get("normalize") or "strip").strip().lower() or "strip" + bv = _norm_expr(f"{before_alias}.data", name, norm) + av = _norm_expr(f"{after_alias}.data", name, norm) + clauses.append(f"({bv} IS DISTINCT FROM {av})") + if not clauses: + return "FALSE" + return "(" + " OR ".join(clauses) + ")" + + +def _key_obj_from_data(data: dict[str, Any] | None, key_fields: list[str]) -> dict[str, Any]: + row = data if isinstance(data, dict) else {} + return {f: row.get(f, "") for f in key_fields} + + +def _changes_from_rows( + before: dict[str, Any] | None, + after: dict[str, Any] | None, + compare_fields: list[str], + rules: dict[str, dict[str, Any]], +) -> dict[str, dict[str, Any]]: + from .compare_rules import explain_diff, values_equal + + out: dict[str, dict[str, Any]] = {} + b = before or {} + a = after or {} + for f in compare_fields: + rule = rules.get(f) + bv = b.get(f, "") + av = a.get(f, "") + if values_equal(bv, av, rule=rule): + continue + entry: dict[str, Any] = {"before": bv, "after": av} + reason = explain_diff(bv, av, rule=rule) + if reason: + entry["reason"] = reason + out[f] = entry + return out + + +def run_sql_sheet_compare( + db: Session, + *, + sheet: Mapping[str, Any], + before_batch_id: str, + after_batch_id: str, + store_unchanged: str = "auto", + on_progress: Callable[[str, int], None] | None = None, +) -> dict[str, Any]: + """Compare one sheet via PostgreSQL TEMP tables + FULL OUTER JOIN. + + Returns the same ``{summary, diffs, mapping_stats}`` shape as ``compare_rows``. + """ + if not _dialect_is_postgres(db): + raise RuntimeError("sql_compare_requires_postgres") + + key_fields = [str(k).strip() for k in (sheet.get("key_fields") or []) if str(k).strip()] + field_rules = list(sheet.get("field_rules") or []) + rules = field_rule_map(field_rules) + compare_fields = effective_compare_fields( + list(sheet.get("compare_fields") or []), + field_rules, + ) + row_filters = list(sheet.get("row_filters") or []) + mid = str(sheet.get("metric_id") or "").strip() + bid_b = str(before_batch_id or "").strip() + bid_a = str(after_batch_id or "").strip() + if not mid or not bid_b or not bid_a or not key_fields: + raise ValueError("sql_compare_missing_args") + + filter_sql, filter_params = compile_row_filters_sql(row_filters) + rk = _rk_sql(key_fields) + tag = uuid4().hex[:8] + tb = f"_netx_cmp_b_{tag}" + ta = f"_netx_cmp_a_{tag}" + + # Raw counts (no row_filters) + raw_b = int( + db.execute( + text( + "SELECT count(*) FROM biz_state_metric_row " + "WHERE batch_id = :bid AND metric_id = :mid" + ), + {"bid": bid_b, "mid": mid}, + ).scalar() + or 0 + ) + if on_progress: + on_progress("before", raw_b) + raw_a = int( + db.execute( + text( + "SELECT count(*) FROM biz_state_metric_row " + "WHERE batch_id = :bid AND metric_id = :mid" + ), + {"bid": bid_a, "mid": mid}, + ).scalar() + or 0 + ) + if on_progress: + on_progress("after", raw_a) + + base_params = {"bid": bid_b, "mid": mid, **filter_params} + # Build TEMP sides + for tname, batch_id in ((tb, bid_b), (ta, bid_a)): + db.execute(text(f"DROP TABLE IF EXISTS {tname}")) + params = {**base_params, "bid": batch_id} + # PRESERVE ROWS: compare progress commits must not drop temps mid-run + db.execute( + text( + f""" + CREATE TEMP TABLE {tname} ON COMMIT PRESERVE ROWS AS + SELECT + id, + ({rk}) AS rk, + data_json AS data, + row_number() OVER ( + PARTITION BY ({rk}) + ORDER BY seq ASC, id ASC + ) AS dup_rn + FROM biz_state_metric_row + WHERE batch_id = :bid + AND metric_id = :mid + AND ({filter_sql}) + """ + ), + params, + ) + db.execute(text(f"CREATE INDEX ON {tname} (rk) WHERE dup_rn = 1")) + + before_n = int( + db.execute(text(f"SELECT count(*) FROM {tb}")).scalar() or 0 + ) + after_n = int( + db.execute(text(f"SELECT count(*) FROM {ta}")).scalar() or 0 + ) + if on_progress: + on_progress("after", after_n) + + changed_pred = _changed_predicate(compare_fields, rules) + kind_expr = f""" + CASE + WHEN b.id IS NULL THEN 'added' + WHEN a.id IS NULL THEN 'removed' + WHEN {changed_pred} THEN 'changed' + ELSE 'unchanged' + END + """ + + # Aggregate primary match kinds (first-wins keys only) + agg_rows = db.execute( + text( + f""" + SELECT {kind_expr} AS kind, count(*)::bigint AS n + FROM (SELECT * FROM {tb} WHERE dup_rn = 1) b + FULL OUTER JOIN (SELECT * FROM {ta} WHERE dup_rn = 1) a + ON b.rk = a.rk + GROUP BY 1 + """ + ) + ).mappings().all() + counts = {str(r["kind"]): int(r["n"] or 0) for r in agg_rows} + added = counts.get("added", 0) + removed = counts.get("removed", 0) + changed = counts.get("changed", 0) + unchanged = counts.get("unchanged", 0) + + dup_b = int( + db.execute(text(f"SELECT count(*) FROM {tb} WHERE dup_rn > 1")).scalar() or 0 + ) + dup_a = int( + db.execute(text(f"SELECT count(*) FROM {ta} WHERE dup_rn > 1")).scalar() or 0 + ) + duplicate = dup_b + dup_a + + # Lazy import — avoid circular import with compare_service + from .compare_service import resolve_unchanged_policy + + policy = resolve_unchanged_policy( + store_unchanged, before_n=before_n, after_n=after_n + ) + diffs: list[dict[str, Any]] = [] + + # Fail + duplicate rows (stream into Python — should be << million) + fail_sql = text( + f""" + SELECT + {kind_expr} AS kind, + b.id AS before_row_id, + a.id AS after_row_id, + b.data AS before_data, + a.data AS after_data, + COALESCE(b.rk, a.rk) AS rk + FROM (SELECT * FROM {tb} WHERE dup_rn = 1) b + FULL OUTER JOIN (SELECT * FROM {ta} WHERE dup_rn = 1) a + ON b.rk = a.rk + WHERE {kind_expr} IN ('added', 'removed', 'changed') + """ + ) + for row in db.execute(fail_sql).mappings(): + kind = str(row["kind"] or "") + before = dict(row["before_data"] or {}) if row["before_data"] is not None else None + after = dict(row["after_data"] or {}) if row["after_data"] is not None else None + key_src = after if kind == "added" else (before or after or {}) + item: dict[str, Any] = { + "kind": kind, + "key": _key_obj_from_data(key_src, key_fields), + "before": before, + "after": after, + "mapped_before": before, + "changes": {}, + "before_row_id": str(row["before_row_id"] or ""), + "after_row_id": str(row["after_row_id"] or ""), + } + if kind == "changed": + item["changes"] = _changes_from_rows(before, after, compare_fields, rules) + diffs.append(item) + + # Duplicate extras + for side, tname in (("before", tb), ("after", ta)): + q = text( + f""" + SELECT id, data, rk FROM {tname} WHERE dup_rn > 1 + """ + ) + for row in db.execute(q).mappings(): + data = dict(row["data"] or {}) + diffs.append( + { + "kind": "duplicate", + "side": side, + "key": _key_obj_from_data(data, key_fields), + "before": data if side == "before" else None, + "after": data if side == "after" else None, + "mapped_before": data if side == "before" else None, + "changes": {}, + "before_row_id": str(row["id"] or "") if side == "before" else "", + "after_row_id": str(row["id"] or "") if side == "after" else "", + } + ) + + unchanged_listed = 0 + include_u = bool(policy.get("include")) + compact = bool(policy.get("compact")) + limit_n = policy.get("limit") + if include_u and unchanged > 0: + lim_sql = "" + params_u: dict[str, Any] = {} + if limit_n is not None: + lim_sql = " LIMIT :lim" + params_u["lim"] = max(0, int(limit_n)) + u_sql = text( + f""" + SELECT + b.id AS before_row_id, + a.id AS after_row_id, + b.data AS before_data, + a.data AS after_data + FROM (SELECT * FROM {tb} WHERE dup_rn = 1) b + INNER JOIN (SELECT * FROM {ta} WHERE dup_rn = 1) a + ON b.rk = a.rk + WHERE NOT ({changed_pred}) + {lim_sql} + """ + ) + for row in db.execute(u_sql, params_u).mappings(): + before = dict(row["before_data"] or {}) + after = dict(row["after_data"] or {}) + unchanged_listed += 1 + if compact: + diffs.append( + { + "kind": "unchanged", + "key": _key_obj_from_data(before, key_fields), + "before": {}, + "after": {}, + "mapped_before": {}, + "changes": {}, + "compact": True, + "before_row_id": str(row["before_row_id"] or ""), + "after_row_id": str(row["after_row_id"] or ""), + } + ) + else: + diffs.append( + { + "kind": "unchanged", + "key": _key_obj_from_data(before, key_fields), + "before": before, + "after": after, + "mapped_before": before, + "changes": {}, + "before_row_id": str(row["before_row_id"] or ""), + "after_row_id": str(row["after_row_id"] or ""), + } + ) + + # Cleanup (also ON COMMIT DROP) + db.execute(text(f"DROP TABLE IF EXISTS {tb}")) + db.execute(text(f"DROP TABLE IF EXISTS {ta}")) + + dup_key_list: list[str] = [] + summary = { + "before_count": before_n, + "after_count": after_n, + "added": added, + "removed": removed, + "changed": changed, + "unchanged": unchanged, + "duplicate": duplicate, + "match_key_fields": list(key_fields), + "duplicate_keys_before": dup_b, + "duplicate_keys_after": dup_a, + "duplicate_key_list": dup_key_list, + "unchanged_listed": unchanged_listed, + "unchanged_truncated": bool( + include_u and limit_n is not None and unchanged > unchanged_listed + ), + "unchanged_compact": bool(compact and unchanged_listed > 0), + "engine": "sql", + "before_raw_count": raw_b, + "after_raw_count": raw_a, + "row_filters": len(row_filters), + "unchanged_policy": policy, + } + mapping_stats = { + "before_iface_count": 0, + "after_iface_count": 0, + "map_pairs": 0, + "hit_before": [], + "miss_before": [], + "hit_after": [], + "miss_after": [], + "unused_before_keys": [], + "ok": True, + "ignore_port_changes": False, + "engine": "sql", + } + return {"summary": summary, "diffs": diffs, "mapping_stats": mapping_stats} diff --git a/tests/test_biz_state_compare_sql.py b/tests/test_biz_state_compare_sql.py new file mode 100644 index 0000000..f885318 --- /dev/null +++ b/tests/test_biz_state_compare_sql.py @@ -0,0 +1,215 @@ +"""Unit tests for SQL compare eligibility and filter compilation.""" + +from __future__ import annotations + +import unittest +from unittest.mock import MagicMock + +from netx_api.biz_state.compare_sql import ( + can_sql_compare, + compile_row_filters_sql, + _field_rules_sql_compatible, + _filters_sql_compatible, + _safe_field, +) + + +class _FakeDialect: + def __init__(self, name: str) -> None: + self.name = name + + +class _FakeBind: + def __init__(self, name: str) -> None: + self.dialect = _FakeDialect(name) + + +def _db(dialect: str = "postgresql") -> MagicMock: + db = MagicMock() + db.get_bind.return_value = _FakeBind(dialect) + return db + + +class CompareSqlGateTests(unittest.TestCase): + def test_safe_field_rejects_injection(self) -> None: + with self.assertRaises(ValueError): + _safe_field("a'; drop table x;--") + with self.assertRaises(ValueError): + _safe_field("x.y") + self.assertEqual(_safe_field("network"), "network") + + def test_requires_postgres(self) -> None: + sheet = { + "metric_id": "bgp_route", + "key_fields": ["network", "next_hop"], + "compare_fields": ["path"], + "iface_fields": [], + "field_rules": [], + "row_filters": [], + } + self.assertFalse(can_sql_compare(_db("sqlite"), sheet, port_map={})) + self.assertTrue(can_sql_compare(_db("postgresql"), sheet, port_map={})) + + def test_rejects_port_map(self) -> None: + sheet = { + "metric_id": "lldp_neighbor", + "key_fields": ["local_if", "remote_sys"], + "compare_fields": [], + "iface_fields": ["local_if"], + "ignore_port_changes": False, + "field_rules": [], + "row_filters": [], + } + self.assertFalse( + can_sql_compare(_db(), sheet, port_map={"gei-0/1": "gei-0/2"}) + ) + + def test_rejects_auto_ignore_ports_with_iface_key(self) -> None: + sheet = { + "metric_id": "lldp_neighbor", + "key_fields": ["local_if", "remote_sys"], + "compare_fields": [], + "iface_fields": ["local_if"], + "ignore_port_changes": None, + "field_rules": [], + "row_filters": [], + } + self.assertFalse(can_sql_compare(_db(), sheet, port_map={})) + + def test_rejects_mac_normalize(self) -> None: + sheet = { + "metric_id": "arp", + "key_fields": ["ip", "vrf"], + "compare_fields": ["mac"], + "iface_fields": [], + "field_rules": [{"field": "mac", "normalize": "mac"}], + "row_filters": [], + } + self.assertFalse(can_sql_compare(_db(), sheet, port_map={})) + + def test_rejects_age_timer_filter(self) -> None: + sheet = { + "metric_id": "arp", + "key_fields": ["ip"], + "compare_fields": [], + "iface_fields": [], + "field_rules": [], + "row_filters": [{"field": "age", "op": "age_timer"}], + } + self.assertFalse(can_sql_compare(_db(), sheet, port_map={})) + + def test_accepts_bgp_route_style(self) -> None: + sheet = { + "metric_id": "bgp_route", + "key_fields": [ + "local_as", + "afi", + "vrf", + "neighbor", + "direction", + "rd", + "network", + "next_hop", + ], + "compare_fields": ["path", "as_num"], + "iface_fields": [], + "field_rules": [], + "row_filters": [{"field": "afi", "op": "eq", "value": "ipv4"}], + } + self.assertTrue(can_sql_compare(_db(), sheet, port_map={})) + + def test_rejects_iface_normalize_when_key_uses_iface(self) -> None: + sheet = { + "metric_id": "interface_brief", + "key_fields": ["interface"], + "compare_fields": ["admin"], + "iface_fields": ["interface"], + "ignore_port_changes": False, + "field_rules": [], + "row_filters": [], + } + self.assertFalse( + can_sql_compare( + _db(), + sheet, + port_map={}, + iface_normalize_rules=[{"from": "GE", "to": "gei"}], + ) + ) + + def test_accepts_ignore_port_changes_false_without_normalize(self) -> None: + sheet = { + "metric_id": "interface_brief", + "key_fields": ["interface"], + "compare_fields": ["admin"], + "iface_fields": ["interface"], + "ignore_port_changes": False, + "field_rules": [], + "row_filters": [], + } + self.assertTrue(can_sql_compare(_db(), sheet, port_map={}, iface_normalize_rules=[])) + + +class CompareSqlFilterCompileTests(unittest.TestCase): + def test_empty_filters(self) -> None: + sql, params = compile_row_filters_sql([]) + self.assertEqual(sql, "TRUE") + self.assertEqual(params, {}) + + def test_eq_case_insensitive(self) -> None: + sql, params = compile_row_filters_sql( + [{"field": "afi", "op": "eq", "value": "IPv4"}] + ) + self.assertIn("lower(", sql) + self.assertIn("afi", sql) + self.assertEqual(list(params.values()), ["ipv4"]) + + def test_contains(self) -> None: + sql, params = compile_row_filters_sql( + [{"field": "af", "op": "contains", "value": "IPv4"}] + ) + self.assertIn("LIKE", sql) + self.assertEqual(list(params.values()), ["%ipv4%"]) + + def test_any_all_nesting(self) -> None: + sql, params = compile_row_filters_sql( + [ + { + "any": [ + {"field": "entry_type", "op": "eq", "value": "dynamic"}, + { + "all": [ + {"field": "entry_type", "op": "empty"}, + {"field": "ip", "op": "not_empty"}, + ] + }, + ] + } + ] + ) + self.assertIn(" OR ", sql) + self.assertIn(" AND ", sql) + self.assertTrue(_filters_sql_compatible([{"any": [{"field": "a", "op": "eq", "value": "1"}]}])) + + def test_in_list(self) -> None: + sql, params = compile_row_filters_sql( + [{"field": "afi", "op": "in", "value": ["ipv4", "ipv6"]}] + ) + self.assertIn(" IN (", sql) + self.assertEqual(sorted(params.values()), ["ipv4", "ipv6"]) + + def test_field_rules_matrix(self) -> None: + self.assertTrue(_field_rules_sql_compatible([{"field": "mac", "normalize": "lower"}])) + self.assertFalse(_field_rules_sql_compatible([{"field": "mac", "normalize": "mac"}])) + self.assertFalse( + _field_rules_sql_compatible( + [{"field": "rx", "compare": "numeric", "tolerance": 1}] + ) + ) + self.assertTrue( + _field_rules_sql_compatible([{"field": "x", "compare": "ignore"}]) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/web/src/pages/network/BizComparePage.tsx b/web/src/pages/network/BizComparePage.tsx index bacf72a..ba8c9ca 100644 --- a/web/src/pages/network/BizComparePage.tsx +++ b/web/src/pages/network/BizComparePage.tsx @@ -138,6 +138,8 @@ type RunSheet = { display_fields?: string[]; field_rules?: FieldRule[]; mode?: string; + /** pending|running|queued|done|cancelled — set while overall run is active */ + status?: string; summary?: Record; diffs?: DiffRow[]; }; @@ -1061,14 +1063,10 @@ export function BizComparePage({ pageMode = "all" }: { pageMode?: BizComparePage useEffect(() => { const runId = String(runDetail?.id || ""); const mid = resultSheetId || sheetIdentity(activeRunSheet) || ""; - const st = String(runDetail?.status || ""); - if ( - !runId || - !mid || - jobDetailTab !== "result" || - st === "running" || - st === "queued" - ) { + const sheetSt = String(activeRunSheet?.status || ""); + // Block only while *this* sheet is still in flight — done sheets are readable mid-run + const sheetStillRunning = ["pending", "running", "queued"].includes(sheetSt); + if (!runId || !mid || jobDetailTab !== "result" || sheetStillRunning) { setPagedDiffs([]); setResultTotal(0); return; @@ -1110,6 +1108,7 @@ export function BizComparePage({ pageMode = "all" }: { pageMode?: BizComparePage runDetail?.status, resultSheetId, activeRunSheet?.metric_id, + activeRunSheet?.status, kindFilter, debouncedResultKw, resultPage, @@ -3714,7 +3713,7 @@ export function BizComparePage({ pageMode = "all" }: { pageMode?: BizComparePage !Number(runDetail?.summary?.unchanged_listed || 0) ? t("bizCompare.unchangedNotStored") : t("bizCompare.resultEmpty") - : runIsActive + : activeSheetPending ? t("bizCompare.runStatusRunning") : t("bizCompare.resultEmpty")}