Add SQL JOIN compare for pushdown-safe sheets and mid-run result viewing.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-10-08 15:16:40 +08:00
parent 958942c5b6
commit ba16167c8e
4 changed files with 1004 additions and 30 deletions

View file

@ -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 {})

View file

@ -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}

View file

@ -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()

View file

@ -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<string, number>;
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")}
</div>