Unify compare search via live source lookup; stratify success samples.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-10-08 17:47:40 +08:00
parent dd4195cdbc
commit b8c0f96982
8 changed files with 546 additions and 40 deletions

View file

@ -2,7 +2,7 @@
from __future__ import annotations
from collections import defaultdict
from collections import defaultdict, deque
from typing import Any, Mapping, Sequence
from .compare_rules import field_rule_map, values_equal, explain_diff
@ -12,6 +12,53 @@ from .iface_normalize import (
resolve_mapped_iface,
)
# Prefer these fields when stratifying success samples (BGP multipath / multi-cmd).
_STRATUM_FIELDS = ("neighbor", "direction", "afi", "vrf")
def stratum_key(row: Mapping[str, Any] | None) -> str:
"""Bucket key for stratified unchanged sampling."""
if not isinstance(row, Mapping):
return "_"
parts: list[str] = []
for f in _STRATUM_FIELDS:
v = str(row.get(f) or "").strip()
if v:
parts.append(f"{f}={v}")
return "|".join(parts) if parts else "_"
def stratify_take(items: list[Any], limit: int, *, key_fn) -> list[Any]:
"""Round-robin across strata so one neighbor/direction cannot consume the whole sample."""
lim = max(0, int(limit))
if lim <= 0 or not items:
return []
if len(items) <= lim:
return list(items)
buckets: dict[str, deque[Any]] = defaultdict(deque)
order: list[str] = []
for it in items:
sk = str(key_fn(it) or "_")
if sk not in buckets:
order.append(sk)
buckets[sk].append(it)
out: list[Any] = []
while len(out) < lim and buckets:
drained: list[str] = []
for sk in order:
q = buckets.get(sk)
if not q:
drained.append(sk)
continue
out.append(q.popleft())
if len(out) >= lim:
break
for sk in drained:
buckets.pop(sk, None)
if sk in order:
order = [x for x in order if x != sk]
return out
def apply_port_map(
row: dict[str, Any],
@ -198,6 +245,7 @@ def compare_rows(
limit_n = None if unchanged_limit is None else max(0, int(unchanged_limit))
multi_before_keys: list[tuple[str, ...]] = []
multi_after_keys: list[tuple[str, ...]] = []
unchanged_candidates: list[tuple[dict[str, Any], dict[str, Any], dict[str, Any]]] = []
def _key_obj(row: dict[str, Any]) -> dict[str, Any]:
return {f: row.get(f, "") for f in key_fields}
@ -210,12 +258,10 @@ def compare_rows(
return str(netx.get("row_id") or "")
return ""
def _emit_unchanged(orig: dict[str, Any], mapped: dict[str, Any], after_row: dict[str, Any]) -> None:
def _append_unchanged_diff(
orig: dict[str, Any], mapped: dict[str, Any], after_row: dict[str, Any]
) -> None:
nonlocal unchanged_listed
if not include_unchanged:
return
if limit_n is not None and unchanged_listed >= limit_n:
return
unchanged_listed += 1
before_rid = _row_id(orig)
after_rid = _row_id(after_row)
@ -330,7 +376,20 @@ def compare_rows(
)
else:
unchanged += 1
_emit_unchanged(orig, mapped, after)
if include_unchanged:
unchanged_candidates.append((orig, mapped, after))
if include_unchanged and unchanged_candidates:
if limit_n is None:
picked = unchanged_candidates
else:
picked = stratify_take(
unchanged_candidates,
limit_n,
key_fn=lambda t: stratum_key(t[1]),
)
for orig, mapped, after_row in picked:
_append_unchanged_diff(orig, mapped, after_row)
def _fmt_keys(keys: list[tuple[str, ...]]) -> list[str]:
seen: set[str] = set()
@ -368,6 +427,11 @@ def compare_rows(
include_unchanged and limit_n is not None and unchanged > unchanged_listed
),
"unchanged_compact": bool(compact_unchanged and unchanged_listed > 0),
"unchanged_sample_mode": (
"stratified"
if include_unchanged and limit_n is not None and unchanged > unchanged_listed
else ("full" if include_unchanged else "none")
),
},
"diffs": diffs,
"mapping_stats": stats,

View file

@ -190,6 +190,9 @@ _SEARCH_TEXT_MAX = 4000
_PERSIST_PROGRESS_EVERY = 10_000
# Above this, store fail diffs as key + row_id + changes (hydrate sides on read).
_FAIL_COMPACT_MIN = 50_000
# Live search (kw): load at most this many matching rows per side, return ≤ this many pairs.
_LIVE_SEARCH_LOAD_CAP = 2_000
_LIVE_SEARCH_RESULT_CAP = 200
# Success-row persist policy (see resolve_unchanged_policy)
_STORE_UNCHANGED_MODES = frozenset({"auto", "always", "never", "sample", "keys"})
_UNCHANGED_FULL_MAX = 20_000
@ -554,6 +557,7 @@ def _pending_sheet_meta(sheet: dict[str, Any]) -> dict[str, Any]:
"compare_fields": compare_fields,
"display_fields": list(sheet.get("display_fields") or []),
"field_rules": list(sheet.get("field_rules") or []),
"row_filters": list(sheet.get("row_filters") or []),
"ignore_port_changes": sheet.get("ignore_port_changes"),
"mode": "presence" if not compare_fields else "fields",
"status": "pending",
@ -2557,6 +2561,7 @@ def _execute_compare_into_run(db: Session, run_id: str) -> dict[str, Any]:
"compare_fields": one["compare_fields"],
"display_fields": one.get("display_fields") or [],
"field_rules": one.get("field_rules") or [],
"row_filters": one.get("row_filters") or [],
"ignore_port_changes": one.get("ignore_port_changes"),
"mode": one["mode"],
"status": "done",
@ -3038,6 +3043,270 @@ def _lookup_sheet(sheets: list[dict[str, Any]], key: str) -> dict[str, Any] | No
return None
def _kind_allows(kind_n: str, diff_kind: str) -> bool:
dk = str(diff_kind or "")
if kind_n == "all":
return True
if kind_n == "diff":
return dk in ("removed", "changed")
return dk == kind_n
def _kw_match_sql(key_fields: list[str], *, param: str = "kw") -> str:
"""OR of ILIKE on key fields (and data_json::text fallback)."""
from .compare_sql import _FIELD_RE, _safe_field
parts: list[str] = []
for f in key_fields:
name = str(f or "").strip()
if not name or not _FIELD_RE.match(name):
continue
sf = _safe_field(name)
parts.append(f"lower(trim(both from coalesce(data_json->>'{sf}', ''))) LIKE :{param}")
# Broad fallback so free-text still hits non-key columns (path, next_hop, …)
parts.append(f"lower(data_json::text) LIKE :{param}")
return "(" + " OR ".join(parts) + ")" if parts else f"(lower(data_json::text) LIKE :{param})"
def _load_metric_rows_for_search(
db: Session,
*,
batch_id: str,
metric_id: str,
row_filters: list[dict[str, Any]] | None,
key_fields: list[str],
kw: str,
cap: int = _LIVE_SEARCH_LOAD_CAP,
) -> tuple[list[dict[str, Any]], bool]:
"""Load rows matching sheet filters + kw. Returns (rows, truncated)."""
from .compare_sql import (
_dialect_is_postgres,
_filters_sql_compatible,
compile_row_filters_sql,
)
from ..models import BizStateMetricRow
from sqlalchemy import text as sql_text
bid = str(batch_id or "").strip()
mid = str(metric_id or "").strip()
needle = str(kw or "").strip()
if not bid or not mid or not needle:
return [], False
lim = max(1, min(int(cap), _LIVE_SEARCH_LOAD_CAP))
filters = [f for f in (row_filters or []) if isinstance(f, dict)]
like = f"%{needle.lower()}%"
if _dialect_is_postgres(db):
filter_sql, filter_params = ("TRUE", {})
if filters and _filters_sql_compatible(filters):
filter_sql, filter_params = compile_row_filters_sql(filters)
kw_sql = _kw_match_sql(key_fields)
# Fetch lim+1 to detect truncation
params = {
"bid": bid,
"mid": mid,
"kw": like,
"lim": lim + 1,
**filter_params,
}
rows = db.execute(
sql_text(
f"""
SELECT id, batch_command_id, task_id, ne_id, seq, data_json, collected_at
FROM biz_state_metric_row
WHERE batch_id = :bid
AND metric_id = :mid
AND ({filter_sql})
AND ({kw_sql})
ORDER BY seq ASC, id ASC
LIMIT :lim
"""
),
params,
).mappings().all()
truncated = len(rows) > lim
rows = rows[:lim]
out: list[dict[str, Any]] = []
for r in rows:
data = dict(r["data_json"] or {})
collected = r["collected_at"]
out.append(
{
**data,
"_netx": {
"batch_id": bid,
"batch_command_id": str(r["batch_command_id"] or ""),
"task_id": str(r["task_id"] or ""),
"ne_id": str(r["ne_id"] or ""),
"collected_at": collected.isoformat() + "Z"
if collected is not None
else None,
"row_id": str(r["id"]),
},
}
)
return out, truncated
# Non-PG / fallback: scan with early stop (OK for tests / small sheets)
q = (
db.query(BizStateMetricRow)
.filter(
BizStateMetricRow.batch_id == bid,
BizStateMetricRow.metric_id == mid,
)
.order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc())
)
out = []
truncated = False
needle_l = needle.lower()
key_set = [str(k).strip() for k in key_fields if str(k).strip()]
for r in q.yield_per(500):
data = dict(r.data_json or {})
row = {
**data,
"_netx": {
"batch_id": bid,
"batch_command_id": r.batch_command_id or "",
"task_id": r.task_id or "",
"ne_id": r.ne_id or "",
"collected_at": r.collected_at.isoformat() + "Z"
if r.collected_at
else None,
"row_id": r.id,
},
}
if filters and not all(row_matches_filter(row, f) for f in filters):
db.expunge(r)
continue
hit = False
for kf in key_set:
if needle_l in str(row.get(kf) or "").lower():
hit = True
break
if not hit:
blob = json.dumps(data, ensure_ascii=False, default=str).lower()
hit = needle_l in blob
if not hit:
db.expunge(r)
continue
out.append(row)
db.expunge(r)
if len(out) >= lim:
# peek one more?
truncated = True
break
return out, truncated
def _live_search_sheet_diffs(
db: Session,
run: BizCompareRun,
sheet: dict[str, Any],
*,
kind: str,
kw: str,
page: int,
page_size: int,
) -> dict[str, Any]:
"""Search before/after metric tables, zip-compare, filter by kind tab."""
mid_src = str(sheet.get("metric_id") or "").strip()
sid = sheet_key(sheet)
key_fields = list(sheet.get("key_fields") or [])
iface_fields = list(sheet.get("iface_fields") or [])
field_rules = list(sheet.get("field_rules") or [])
compare_fields = effective_compare_fields(
list(sheet.get("compare_fields") or []),
field_rules,
)
row_filters = list(sheet.get("row_filters") or [])
# Older runs may lack row_filters on sheet meta — fall back to template
if not row_filters and run.template_id:
tpl = db.get(BizCompareTemplate, run.template_id)
if tpl:
for s in template_metrics(tpl):
if sheet_key(s) == sid:
row_filters = list(s.get("row_filters") or [])
if not key_fields:
key_fields = list(s.get("key_fields") or [])
if not field_rules:
field_rules = list(s.get("field_rules") or [])
compare_fields = effective_compare_fields(
list(s.get("compare_fields") or compare_fields),
field_rules,
)
break
if not key_fields or not mid_src:
return {
"total": 0,
"page": page,
"page_size": page_size,
"metric_id": sid,
"items": [],
"source": "live",
"truncated": False,
}
ignore_ports = sheet.get("ignore_port_changes")
if ignore_ports is not None:
ignore_ports = bool(ignore_ports)
pmap = _port_map_dict(db, str(run.mapping_id or ""))
tpl = db.get(BizCompareTemplate, run.template_id) if run.template_id else None
norm_rules = template_iface_normalize(tpl)
before_rows, trunc_b = _load_metric_rows_for_search(
db,
batch_id=str(run.before_batch_id or ""),
metric_id=mid_src,
row_filters=row_filters,
key_fields=key_fields,
kw=kw,
)
after_rows, trunc_a = _load_metric_rows_for_search(
db,
batch_id=str(run.after_batch_id or ""),
metric_id=mid_src,
row_filters=row_filters,
key_fields=key_fields,
kw=kw,
)
result = compare_rows(
before_rows=before_rows,
after_rows=after_rows,
key_fields=key_fields,
iface_fields=iface_fields,
compare_fields=compare_fields,
port_map=pmap,
field_rules=field_rules,
iface_normalize_rules=norm_rules,
ignore_port_changes=ignore_ports,
include_unchanged=True,
unchanged_limit=None,
compact_unchanged=False,
)
kind_n = (kind or "diff").strip().lower()
filtered = [
d for d in list(result.get("diffs") or []) if _kind_allows(kind_n, str(d.get("kind") or ""))
]
# Cap pairs returned to keep UI snappy
truncated = bool(trunc_b or trunc_a or len(filtered) > _LIVE_SEARCH_RESULT_CAP)
filtered = filtered[:_LIVE_SEARCH_RESULT_CAP]
total = len(filtered)
start = (page - 1) * page_size
page_items = filtered[start : start + page_size]
return {
"total": total,
"page": page,
"page_size": page_size,
"metric_id": sid,
"items": page_items,
"source": "live",
"truncated": truncated,
"live_before_matched": len(before_rows),
"live_after_matched": len(after_rows),
}
def list_run_diffs(
db: Session,
run_id: str,
@ -3063,6 +3332,18 @@ def list_run_diffs(
sheet = _lookup_sheet(sheets, asked) if asked else (sheets[0] if sheets else None)
mid = sheet_key(sheet) if sheet else (asked or str(r.metric_id or ""))
# Unified search: any kw → live source lookup; kind tab only filters.
if kw_n and sheet:
return _live_search_sheet_diffs(
db,
r,
sheet,
kind=kind_n,
kw=kw_n,
page=page_n,
page_size=size_n,
)
if _run_has_diff_rows(db, run_id):
q = db.query(BizCompareDiff).filter(
BizCompareDiff.run_id == run_id,
@ -3070,10 +3351,12 @@ def list_run_diffs(
)
if kind_n == "diff":
q = q.filter(BizCompareDiff.kind.in_(("removed", "changed")))
elif kind_n == "removed":
q = q.filter(BizCompareDiff.kind == "removed")
elif kind_n == "changed":
q = q.filter(BizCompareDiff.kind == "changed")
elif kind_n != "all":
q = q.filter(BizCompareDiff.kind == kind_n)
if kw_n:
q = q.filter(BizCompareDiff.search_text.ilike(f"%{kw_n}%"))
total = q.count()
rows = (
q.order_by(BizCompareDiff.seq.asc(), BizCompareDiff.id.asc())
@ -3088,6 +3371,8 @@ def list_run_diffs(
"page_size": size_n,
"metric_id": mid,
"items": items,
"source": "stored",
"truncated": False,
}
# Legacy: diffs embedded in summary_json / diffs_json
@ -3097,7 +3382,7 @@ def list_run_diffs(
inline = list((sheet or {}).get("diffs") or [])
if not inline and mid == r.metric_id:
inline = list(r.diffs_json or [])
filtered = _filter_inline_diffs(inline, kind=kind_n, kw=kw_n)
filtered = _filter_inline_diffs(inline, kind=kind_n, kw="")
total = len(filtered)
start = (page_n - 1) * size_n
page_items = filtered[start : start + size_n]
@ -3107,6 +3392,8 @@ def list_run_diffs(
"page_size": size_n,
"metric_id": mid,
"items": page_items,
"source": "stored",
"truncated": False,
}

View file

@ -613,25 +613,66 @@ def run_sql_sheet_compare(
compact = bool(policy.get("compact"))
limit_n = policy.get("limit")
if include_u and unchanged > 0:
lim_sql = ""
params_u: dict[str, Any] = {}
# Stratified sample: round-robin by neighbor|direction|afi|vrf so one
# BGP command cannot consume the whole success sample.
stratum_expr = """
concat_ws('|',
coalesce(nullif(trim(both from coalesce(b.data->>'neighbor','')), ''), '_'),
coalesce(nullif(trim(both from coalesce(b.data->>'direction','')), ''), '_'),
coalesce(nullif(trim(both from coalesce(b.data->>'afi','')), ''), '_'),
coalesce(nullif(trim(both from coalesce(b.data->>'vrf','')), ''), '_')
)
"""
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 {tb} b
INNER JOIN {ta} a
ON b.rk = a.rk AND b.dup_rn = a.dup_rn
WHERE NOT ({changed_pred})
{lim_sql}
"""
)
u_sql = text(
f"""
WITH u AS (
SELECT
b.id AS before_row_id,
a.id AS after_row_id,
b.data AS before_data,
a.data AS after_data,
{stratum_expr} AS stratum
FROM {tb} b
INNER JOIN {ta} a
ON b.rk = a.rk AND b.dup_rn = a.dup_rn
WHERE NOT ({changed_pred})
),
ranked AS (
SELECT *,
row_number() OVER (
PARTITION BY stratum ORDER BY before_row_id ASC
) AS rn_in
FROM u
),
picked AS (
SELECT *,
row_number() OVER (
ORDER BY rn_in ASC, stratum ASC, before_row_id ASC
) AS pick_ord
FROM ranked
)
SELECT before_row_id, after_row_id, before_data, after_data
FROM picked
WHERE pick_ord <= :lim
"""
)
else:
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 {tb} b
INNER JOIN {ta} a
ON b.rk = a.rk AND b.dup_rn = a.dup_rn
WHERE NOT ({changed_pred})
"""
)
for row in db.execute(u_sql, params_u).mappings():
before = dict(row["before_data"] or {})
after = dict(row["after_data"] or {})
@ -685,6 +726,11 @@ def run_sql_sheet_compare(
include_u and limit_n is not None and unchanged > unchanged_listed
),
"unchanged_compact": bool(compact and unchanged_listed > 0),
"unchanged_sample_mode": (
"stratified"
if include_u and limit_n is not None and unchanged > unchanged_listed
else ("full" if include_u else "none")
),
"engine": "sql",
"before_raw_count": raw_b,
"after_raw_count": raw_a,