fix(biz-compare): validate templates and protect comparison inputs

This commit is contained in:
oliver 2026-10-11 05:56:09 +08:00
parent 8c15975d3e
commit 90cf1907b4
7 changed files with 376 additions and 54 deletions

View file

@ -7,7 +7,7 @@ from heapq import heappush, heapreplace
from itertools import chain from itertools import chain
from typing import Any, Iterable, Mapping, Sequence from typing import Any, Iterable, Mapping, Sequence
from .compare_rules import field_rule_map, values_equal, explain_diff from .compare_rules import field_rule_map, values_equal, explain_diff, scalar_text
from .iface_normalize import ( from .iface_normalize import (
apply_iface_normalize_rows, apply_iface_normalize_rows,
normalize_iface_rules, normalize_iface_rules,
@ -94,7 +94,7 @@ def apply_port_map(
def row_key(row: dict[str, Any], key_fields: list[str]) -> tuple[str, ...]: def row_key(row: dict[str, Any], key_fields: list[str]) -> tuple[str, ...]:
return tuple(str(row.get(k) or "").strip() for k in key_fields) return tuple(scalar_text(row.get(k)).strip() for k in key_fields)
def mapping_stats( def mapping_stats(

View file

@ -7,10 +7,21 @@ All metric-specific compare behavior belongs in the sheet template
from __future__ import annotations from __future__ import annotations
import re import re
import math
from typing import Any, Mapping, Sequence from typing import Any, Mapping, Sequence
_AGE_TIMER_RE = re.compile(r"^\d{1,2}:\d{2}:\d{2}$") _AGE_TIMER_RE = re.compile(r"^\d{1,2}:\d{2}:\d{2}$")
NUMBER_PATTERN = r"^[+-]?([0-9]+(\.[0-9]*)?|\.[0-9]+)([eE][+-]?[0-9]+)?$"
def scalar_text(value: Any) -> str:
"""Match JSON text extraction: only null is blank; keep zero and booleans."""
if value is None:
return ""
if isinstance(value, bool):
return "true" if value else "false"
return str(value)
def _as_list(raw: Any) -> list[Any]: def _as_list(raw: Any) -> list[Any]:
@ -22,7 +33,7 @@ def _as_list(raw: Any) -> list[Any]:
def _field_val(row: Mapping[str, Any], field: str) -> str: def _field_val(row: Mapping[str, Any], field: str) -> str:
return str((row or {}).get(field) or "").strip() return scalar_text((row or {}).get(field)).strip()
def eval_leaf_filter(row: Mapping[str, Any], filt: Mapping[str, Any]) -> bool: def eval_leaf_filter(row: Mapping[str, Any], filt: Mapping[str, Any]) -> bool:
@ -35,17 +46,17 @@ def eval_leaf_filter(row: Mapping[str, Any], filt: Mapping[str, Any]) -> bool:
expect = filt.get("value") expect = filt.get("value")
if op in ("eq", "=="): if op in ("eq", "=="):
return raw.lower() == str(expect or "").strip().lower() return raw.lower() == scalar_text(expect).strip().lower()
if op in ("ne", "!="): if op in ("ne", "!="):
return raw.lower() != str(expect or "").strip().lower() return raw.lower() != scalar_text(expect).strip().lower()
if op == "in": if op == "in":
opts = {str(x).strip().lower() for x in _as_list(expect) if str(x).strip()} opts = {scalar_text(x).strip().lower() for x in _as_list(expect) if scalar_text(x).strip()}
return raw.lower() in opts return raw.lower() in opts
if op in ("not_in", "nin"): if op in ("not_in", "nin"):
opts = {str(x).strip().lower() for x in _as_list(expect) if str(x).strip()} opts = {scalar_text(x).strip().lower() for x in _as_list(expect) if scalar_text(x).strip()}
return raw.lower() not in opts return raw.lower() not in opts
if op == "contains": if op == "contains":
needle = str(expect or "").strip().lower() needle = scalar_text(expect).strip().lower()
return bool(needle) and needle in raw.lower() return bool(needle) and needle in raw.lower()
if op == "empty": if op == "empty":
return not raw return not raw
@ -63,7 +74,7 @@ def eval_leaf_filter(row: Mapping[str, Any], filt: Mapping[str, Any]) -> bool:
# HH:MM:SS dynamic ARP age # HH:MM:SS dynamic ARP age
return bool(_AGE_TIMER_RE.match(raw)) return bool(_AGE_TIMER_RE.match(raw))
if op == "ci_eq": if op == "ci_eq":
return raw.lower() == str(expect or "").strip().lower() return raw.lower() == scalar_text(expect).strip().lower()
return True return True
@ -99,7 +110,7 @@ def apply_row_filters(
def normalize_value(value: Any, how: str) -> str: def normalize_value(value: Any, how: str) -> str:
text = str(value if value is not None else "").strip() text = scalar_text(value).strip()
mode = str(how or "").strip().lower() mode = str(how or "").strip().lower()
if not mode or mode in ("none", "strip"): if not mode or mode in ("none", "strip"):
return text return text
@ -190,8 +201,11 @@ def effective_display_fields(
def _parse_float(text: str) -> float | None: def _parse_float(text: str) -> float | None:
if not re.fullmatch(NUMBER_PATTERN, text):
return None
try: try:
return float(text) if text else 0.0 value = float(text)
return value if math.isfinite(value) else None
except ValueError: except ValueError:
return None return None

View file

@ -45,6 +45,7 @@ from .iface_normalize import (
normalize_iface_rules, normalize_iface_rules,
) )
from .profiles import metric_field_map from .profiles import metric_field_map
from .compare_validation import validate_template_body
_log = logging.getLogger("netx.biz_state.compare") _log = logging.getLogger("netx.biz_state.compare")
@ -92,7 +93,7 @@ def batch_metric_collect_ok(db: Session, batch_id: str, metric_id: str) -> bool:
if status == "success": if status == "success":
return True return True
cmds = ( cmds = (
db.query(BizStateBatchCommand) db.query(BizStateBatchCommand.parse_status, BizStateBatchCommand.row_count)
.filter( .filter(
BizStateBatchCommand.batch_id == bid, BizStateBatchCommand.batch_id == bid,
BizStateBatchCommand.metric_id == mid, BizStateBatchCommand.metric_id == mid,
@ -582,34 +583,6 @@ 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( def _order_sheets_small_first(
db: Session, db: Session,
sheets: list[dict[str, Any]], sheets: list[dict[str, Any]],
@ -620,10 +593,19 @@ def _order_sheets_small_first(
"""Run smaller metrics first so field engineers can review early results.""" """Run smaller metrics first so field engineers can review early results."""
if len(sheets) <= 1: if len(sheets) <= 1:
return list(sheets) return list(sheets)
batches = [before_batch_id, after_batch_id] # Split sheets share metrics. Aggregate all counters once, without raw CLI blobs.
counts: dict[str, int] = {}
for bid, mid, total in db.query(
BizStateBatchCommand.batch_id, BizStateBatchCommand.metric_id,
func.sum(BizStateBatchCommand.row_count),
).filter(
BizStateBatchCommand.batch_id.in_([before_batch_id, after_batch_id]),
BizStateBatchCommand.metric_id.in_({str(s.get("metric_id") or "") for s in sheets}),
).group_by(BizStateBatchCommand.batch_id, BizStateBatchCommand.metric_id):
counts[mid] = max(counts.get(mid, 0), int(total or 0))
scored: list[tuple[int, int, dict[str, Any]]] = [] scored: list[tuple[int, int, dict[str, Any]]] = []
for i, sheet in enumerate(sheets): for i, sheet in enumerate(sheets):
n = _metric_row_estimate(db, batches, str(sheet.get("metric_id") or "")) n = counts.get(str(sheet.get("metric_id") or ""), 0)
scored.append((n, i, sheet)) scored.append((n, i, sheet))
scored.sort(key=lambda x: (x[0], x[1])) scored.sort(key=lambda x: (x[0], x[1]))
return [s for _, _, s in scored] return [s for _, _, s in scored]
@ -1132,6 +1114,7 @@ def _parse_metrics_body(body: dict[str, Any]) -> list[dict[str, Any]]:
display_fields=disp_arg, display_fields=disp_arg,
row_filters=_normalize_row_filters(body.get("row_filters")), row_filters=_normalize_row_filters(body.get("row_filters")),
field_rules=rules, field_rules=rules,
ignore_port_changes=body.get("ignore_port_changes"),
) )
] ]
@ -1405,7 +1388,8 @@ def list_templates(db: Session) -> list[dict[str, Any]]:
ensure_default_templates(db) ensure_default_templates(db)
upgrade_builtin_split_sheets(db) upgrade_builtin_split_sheets(db)
rows = db.query(BizCompareTemplate).order_by(BizCompareTemplate.name.asc()).all() rows = db.query(BizCompareTemplate).order_by(BizCompareTemplate.name.asc()).all()
return [_template_out(t) for t in rows] usage = dict(db.query(BizCompareJob.template_id, func.count(BizCompareJob.id)).group_by(BizCompareJob.template_id).all())
return [{**_template_out(t), "job_count": usage.get(t.id, 0)} for t in rows]
def list_metric_schemas() -> list[dict[str, Any]]: def list_metric_schemas() -> list[dict[str, Any]]:
@ -1442,6 +1426,7 @@ def list_row_filter_presets() -> list[dict[str, Any]]:
def create_template(db: Session, body: dict[str, Any]) -> dict[str, Any]: def create_template(db: Session, body: dict[str, Any]) -> dict[str, Any]:
validate_template_body(body)
sheets = _parse_metrics_body(body) sheets = _parse_metrics_body(body)
t = BizCompareTemplate( t = BizCompareTemplate(
id=uuid4().hex, id=uuid4().hex,
@ -1464,6 +1449,7 @@ def update_template(db: Session, template_id: str, body: dict[str, Any]) -> dict
t = db.get(BizCompareTemplate, template_id) t = db.get(BizCompareTemplate, template_id)
if not t: if not t:
raise HTTPException(status_code=404, detail="template_not_found") raise HTTPException(status_code=404, detail="template_not_found")
validate_template_body(body, partial=True)
if "name" in body: if "name" in body:
t.name = str(body.get("name") or "")[:256] t.name = str(body.get("name") or "")[:256]
if "note" in body: if "note" in body:
@ -1485,6 +1471,7 @@ def update_template(db: Session, template_id: str, body: dict[str, Any]) -> dict
"display_fields", "display_fields",
"row_filters", "row_filters",
"field_rules", "field_rules",
"ignore_port_changes",
) )
): ):
# Prefer explicit metrics; otherwise merge into current sheets from legacy keys # Prefer explicit metrics; otherwise merge into current sheets from legacy keys
@ -1522,6 +1509,8 @@ def update_template(db: Session, template_id: str, body: dict[str, Any]) -> dict
first["row_filters"] = _normalize_row_filters(body.get("row_filters")) first["row_filters"] = _normalize_row_filters(body.get("row_filters"))
if "field_rules" in body: if "field_rules" in body:
first["field_rules"] = _normalize_field_rules(body.get("field_rules")) first["field_rules"] = _normalize_field_rules(body.get("field_rules"))
if "ignore_port_changes" in body:
first["ignore_port_changes"] = body["ignore_port_changes"]
sheets[0] = _normalize_sheet(first) or first sheets[0] = _normalize_sheet(first) or first
_apply_sheets_to_row(t, sheets) _apply_sheets_to_row(t, sheets)
t.updated_at = _utcnow() t.updated_at = _utcnow()
@ -1533,6 +1522,8 @@ def delete_template(db: Session, template_id: str) -> None:
t = db.get(BizCompareTemplate, template_id) t = db.get(BizCompareTemplate, template_id)
if not t: if not t:
raise HTTPException(status_code=404, detail="template_not_found") raise HTTPException(status_code=404, detail="template_not_found")
if db.query(BizCompareJob.id).filter(BizCompareJob.template_id == template_id).first():
raise HTTPException(status_code=409, detail="template_in_use")
db.delete(t) db.delete(t)
db.commit() db.commit()
@ -1888,6 +1879,16 @@ def list_jobs(db: Session) -> list[dict[str, Any]]:
return [_job_out(j) for j in rows] return [_job_out(j) for j in rows]
def _validate_job_sheets(db: Session, template_id: str, enabled: list[str]) -> None:
tpl = db.get(BizCompareTemplate, template_id)
if not tpl:
raise HTTPException(status_code=404, detail="template_not_found")
sheets = {sheet_key(s) for s in template_metrics(tpl)}
unknown = sorted(set(enabled) - sheets)
if unknown:
raise HTTPException(status_code=400, detail={"error": "unknown_enabled_sheets", "sheet_ids": unknown})
def create_job(db: Session, body: dict[str, Any]) -> dict[str, Any]: def create_job(db: Session, body: dict[str, Any]) -> dict[str, Any]:
ensure_default_templates(db) ensure_default_templates(db)
template_id = str(body.get("template_id") or "").strip() template_id = str(body.get("template_id") or "").strip()
@ -1898,6 +1899,7 @@ def create_job(db: Session, body: dict[str, Any]) -> dict[str, Any]:
if not db.get(BizCompareTemplate, template_id): if not db.get(BizCompareTemplate, template_id):
raise HTTPException(status_code=404, detail="template_not_found") raise HTTPException(status_code=404, detail="template_not_found")
enabled = _str_list(body.get("enabled_sheet_ids")) enabled = _str_list(body.get("enabled_sheet_ids"))
_validate_job_sheets(db, template_id, enabled)
j = BizCompareJob( j = BizCompareJob(
id=uuid4().hex, id=uuid4().hex,
name=str(body.get("name") or "compare")[:256], name=str(body.get("name") or "compare")[:256],
@ -1926,6 +1928,8 @@ def update_job(db: Session, job_id: str, body: dict[str, Any]) -> dict[str, Any]
j = db.get(BizCompareJob, job_id) j = db.get(BizCompareJob, job_id)
if not j: if not j:
raise HTTPException(status_code=404, detail="job_not_found") raise HTTPException(status_code=404, detail="job_not_found")
_validate_job_sheets(db, str(body.get("template_id", j.template_id) or ""),
_str_list(body.get("enabled_sheet_ids", j.enabled_sheet_ids)))
for key in ( for key in (
"name", "name",
"template_id", "template_id",
@ -1953,6 +1957,9 @@ def delete_job(db: Session, job_id: str) -> None:
j = db.get(BizCompareJob, job_id) j = db.get(BizCompareJob, job_id)
if not j: if not j:
raise HTTPException(status_code=404, detail="job_not_found") raise HTTPException(status_code=404, detail="job_not_found")
if db.query(BizCompareRun.id).filter(BizCompareRun.job_id == job_id,
BizCompareRun.status.in_(("queued", "running"))).first():
raise HTTPException(status_code=409, detail="compare_running")
run_ids = [ run_ids = [
rid for (rid,) in db.query(BizCompareRun.id).filter(BizCompareRun.job_id == job_id).all() rid for (rid,) in db.query(BizCompareRun.id).filter(BizCompareRun.job_id == job_id).all()
] ]
@ -1970,6 +1977,8 @@ def delete_run(db: Session, run_id: str) -> dict[str, Any]:
r = db.get(BizCompareRun, run_id) r = db.get(BizCompareRun, run_id)
if not r: if not r:
raise HTTPException(status_code=404, detail="run_not_found") raise HTTPException(status_code=404, detail="run_not_found")
if r.status in ("queued", "running"):
raise HTTPException(status_code=409, detail="compare_running")
job_id = str(r.job_id or "") job_id = str(r.job_id or "")
db.query(BizCompareDiff).filter(BizCompareDiff.run_id == run_id).delete( db.query(BizCompareDiff).filter(BizCompareDiff.run_id == run_id).delete(
synchronize_session=False synchronize_session=False
@ -2222,8 +2231,11 @@ def _validate_compare_job(
after_batch_id = str(force_after_batch_id or "").strip() or _resolve_after_batch(db, j) after_batch_id = str(force_after_batch_id or "").strip() or _resolve_after_batch(db, j)
if not before_batch_id or not after_batch_id: if not before_batch_id or not after_batch_id:
raise HTTPException(status_code=400, detail="before_and_after_batch_required") raise HTTPException(status_code=400, detail="before_and_after_batch_required")
if not db.get(BizStateBatch, before_batch_id) or not db.get(BizStateBatch, after_batch_id): source_batches = [db.get(BizStateBatch, bid) for bid in (before_batch_id, after_batch_id)]
if not all(source_batches):
raise HTTPException(status_code=404, detail="batch_not_found") raise HTTPException(status_code=404, detail="batch_not_found")
if any(b.status != "success" for b in source_batches):
raise HTTPException(status_code=409, detail="source_batch_not_complete")
sheets_cfg = template_metrics(tpl) sheets_cfg = template_metrics(tpl)
if not sheets_cfg: if not sheets_cfg:

View file

@ -14,7 +14,7 @@ from uuid import uuid4
from sqlalchemy import text from sqlalchemy import text
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .compare_rules import effective_compare_fields, field_rule_map from .compare_rules import effective_compare_fields, field_rule_map, NUMBER_PATTERN, scalar_text
_log = logging.getLogger("netx.biz_state.compare_sql") _log = logging.getLogger("netx.biz_state.compare_sql")
@ -53,7 +53,7 @@ _SQL_COMPARE_MODES = frozenset(
} }
) )
_EMPTY_AS_BLANK = ("n/a", "na", "-", "--", "none", "null") _EMPTY_AS_BLANK = ("n/a", "na", "-", "--", "none", "null")
_NUM_RE_SQL = r"^-?[0-9]+(\.[0-9]+)?([eE][-+]?[0-9]+)?$" _NUM_RE_SQL = NUMBER_PATTERN
def _dialect_is_postgres(db: Session) -> bool: def _dialect_is_postgres(db: Session) -> bool:
@ -258,27 +258,30 @@ def compile_row_filters_sql(
expect = filt.get("value") expect = filt.get("value")
if op in ("eq", "==", "ci_eq"): if op in ("eq", "==", "ci_eq"):
key = _next("v") key = _next("v")
params[key] = str(expect or "").strip() params[key] = scalar_text(expect).strip()
if op == "ci_eq" or op in ("eq", "=="): if op == "ci_eq" or op in ("eq", "=="):
# Python eq is case-insensitive via .lower() # Python eq is case-insensitive via .lower()
params[key] = str(expect or "").strip().lower() params[key] = scalar_text(expect).strip().lower()
return f"(lower({expr}) = :{key})" return f"(lower({expr}) = :{key})"
if op in ("ne", "!="): if op in ("ne", "!="):
key = _next("v") key = _next("v")
params[key] = str(expect or "").strip().lower() params[key] = scalar_text(expect).strip().lower()
return f"(lower({expr}) <> :{key})" return f"(lower({expr}) <> :{key})"
if op == "contains": if op == "contains":
key = _next("v") key = _next("v")
needle = str(expect or "").strip().lower() needle = scalar_text(expect).strip().lower()
params[key] = f"%{needle}%" if not needle:
return f"(lower({expr}) LIKE :{key})" return "FALSE"
escaped = needle.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
params[key] = f"%{escaped}%"
return f"(lower({expr}) LIKE :{key} ESCAPE E'\\\\')"
if op in ("in", "not_in", "nin"): if op in ("in", "not_in", "nin"):
vals = expect vals = expect
if vals is None: if vals is None:
vals = [] vals = []
if not isinstance(vals, (list, tuple)): if not isinstance(vals, (list, tuple)):
vals = [vals] vals = [vals]
clean = [str(x).strip().lower() for x in vals if str(x).strip()] clean = [scalar_text(x).strip().lower() for x in vals if scalar_text(x).strip()]
if not clean: if not clean:
return "FALSE" if op == "in" else "TRUE" return "FALSE" if op == "in" else "TRUE"
keys = [] keys = []
@ -304,7 +307,8 @@ def _rk_sql(key_fields: list[str], *, json_col: str = "data_json") -> str:
parts.append(f"trim(both from coalesce({json_col}->>'{f}', ''))") parts.append(f"trim(both from coalesce({json_col}->>'{f}', ''))")
if len(parts) == 1: if len(parts) == 1:
return parts[0] return parts[0]
return "concat_ws('|', " + ", ".join(parts) + ")" # Delimiter concatenation conflates ('a|b', 'c') with ('a', 'b|c').
return "jsonb_build_array(" + ", ".join(parts) + ")::text"
def _field_differs_sql(bv: str, av: str, rule: Mapping[str, Any]) -> str: def _field_differs_sql(bv: str, av: str, rule: Mapping[str, Any]) -> str:

View file

@ -0,0 +1,130 @@
"""Validate template writes without silently dropping parts of the comparison."""
from __future__ import annotations
import math
import re
from typing import Any
from fastapi import HTTPException
FILTER_OPS = {"eq", "==", "ci_eq", "ne", "!=", "in", "not_in", "nin", "contains", "empty", "not_empty", "nonempty", "regex", "age_timer"}
COMPARE_MODES = {"", "eq", "ignore", "skip", "off", "numeric", "number", "int", "float", "percent", "pct", "rel"}
NORMALIZE_MODES = {"", "none", "strip", "lower", "upper", "mac", "empty_as_blank"}
def invalid(path: str, reason: str) -> None:
raise HTTPException(status_code=400, detail={"error": "invalid_template", "path": path, "reason": reason})
def _strings(value: Any, path: str, *, required: bool = False) -> None:
if not isinstance(value, list) or any(not isinstance(x, str) or not x.strip() for x in value):
invalid(path, "nonempty_strings_required")
if required and not value:
invalid(path, "key_fields_required")
if len({x.strip() for x in value}) != len(value):
invalid(path, "duplicate_field")
def validate_filters(value: Any, path: str, depth: int = 0) -> None:
if not isinstance(value, list):
invalid(path, "list_required")
if depth > 12:
invalid(path, "filter_nesting_too_deep")
for i, node in enumerate(value):
here = f"{path}[{i}]"
if not isinstance(node, dict) or not node:
invalid(here, "filter_required")
groups = [key for key in ("any", "all") if key in node]
if groups:
if len(groups) != 1 or any(node.get(key) is not None for key in ("field", "op", "value")):
invalid(here, "group_or_leaf_required")
kids = node[groups[0]]
if not isinstance(kids, list) or not kids:
invalid(here, "nonempty_group_required")
validate_filters(kids, f"{here}.{groups[0]}", depth + 1)
continue
if not isinstance(node.get("field"), str) or not node["field"].strip():
invalid(here, "filter_field_required")
op = str(node.get("op") or "eq").strip().lower()
if op not in FILTER_OPS:
invalid(here, "unknown_filter_operator")
expect = node.get("value")
if op in ("in", "not_in", "nin"):
if not isinstance(expect, list) or not expect or any(not isinstance(x, (str, int, float, bool)) for x in expect):
invalid(here, "nonempty_value_list_required")
elif op not in ("empty", "not_empty", "nonempty", "age_timer"):
if not isinstance(expect, (str, int, float, bool)):
invalid(here, "scalar_value_required")
if op in ("contains", "regex") and not str(expect).strip():
invalid(here, "value_required")
if op == "regex":
try:
re.compile(str(expect), re.I)
except re.error:
invalid(here, "invalid_regex")
def validate_sheet(sheet: Any, path: str, *, partial: bool = False) -> None:
if not isinstance(sheet, dict):
invalid(path, "sheet_required")
if not partial or "metric_id" in sheet:
if not isinstance(sheet.get("metric_id"), str) or not sheet["metric_id"].strip():
invalid(path, "metric_id_required")
for key in ("key_fields", "iface_fields", "compare_fields", "display_fields", "ignore_fields"):
if key in sheet or (key == "key_fields" and not partial):
_strings(sheet.get(key), f"{path}.{key}", required=key == "key_fields")
if "ignore_port_changes" in sheet and sheet["ignore_port_changes"] is not None and not isinstance(sheet["ignore_port_changes"], bool):
invalid(path, "boolean_port_policy_required")
if "row_filters" in sheet:
validate_filters(sheet["row_filters"], f"{path}.row_filters")
if "field_rules" in sheet:
rules = sheet["field_rules"]
if not isinstance(rules, list):
invalid(path, "field_rules_list_required")
seen = set()
for i, rule in enumerate(rules):
here = f"{path}.field_rules[{i}]"
if not isinstance(rule, dict) or not isinstance(rule.get("field"), str) or not rule["field"].strip():
invalid(here, "rule_field_required")
field = rule["field"].strip()
if field in seen:
invalid(here, "duplicate_field_rule")
seen.add(field)
if str(rule.get("compare") or "").strip().lower() not in COMPARE_MODES:
invalid(here, "unknown_compare_mode")
if str(rule.get("normalize") or "").strip().lower() not in NORMALIZE_MODES:
invalid(here, "unknown_normalize_mode")
tol = rule.get("tolerance")
if tol is not None:
if isinstance(tol, bool) or not isinstance(tol, (int, float)) or not math.isfinite(tol) or tol < 0:
invalid(here, "finite_nonnegative_tolerance_required")
def validate_template_body(body: dict[str, Any], *, partial: bool = False) -> None:
metrics = body.get("metrics")
if metrics is not None:
if not isinstance(metrics, list) or not metrics:
invalid("metrics", "metrics_required")
seen = set()
for i, sheet in enumerate(metrics):
path = f"metrics[{i}]"
validate_sheet(sheet, path)
sid = str(sheet.get("sheet_id") or sheet["metric_id"]).strip()
if sid in seen:
invalid(path, "duplicate_sheet_id")
seen.add(sid)
else:
validate_sheet(body, "template", partial=partial)
if "iface_normalize_rules" in body and body["iface_normalize_rules"] is not None:
rules = body["iface_normalize_rules"]
if not isinstance(rules, list):
invalid("iface_normalize_rules", "list_required")
seen = set()
for i, rule in enumerate(rules):
path = f"iface_normalize_rules[{i}]"
if not isinstance(rule, dict) or any(not isinstance(rule.get(k), str) or not rule[k].strip() for k in ("from", "to")):
invalid(path, "alias_pair_required")
key = rule["from"].strip().lower()
if key in seen:
invalid(path, "duplicate_alias")
seen.add(key)

View file

@ -574,6 +574,7 @@ class TemplateMetricIn(BaseModel):
display_fields: list[str] = Field(default_factory=list) display_fields: list[str] = Field(default_factory=list)
row_filters: list[dict[str, Any]] = Field(default_factory=list) row_filters: list[dict[str, Any]] = Field(default_factory=list)
field_rules: list[FieldRuleIn] = Field(default_factory=list) field_rules: list[FieldRuleIn] = Field(default_factory=list)
ignore_port_changes: bool | None = None
class TemplateIn(BaseModel): class TemplateIn(BaseModel):
@ -592,6 +593,7 @@ class TemplateIn(BaseModel):
display_fields: list[str] = Field(default_factory=list) display_fields: list[str] = Field(default_factory=list)
row_filters: list[dict[str, Any]] = Field(default_factory=list) row_filters: list[dict[str, Any]] = Field(default_factory=list)
field_rules: list[FieldRuleIn] = Field(default_factory=list) field_rules: list[FieldRuleIn] = Field(default_factory=list)
ignore_port_changes: bool | None = None
class TemplatePatchIn(BaseModel): class TemplatePatchIn(BaseModel):
@ -607,6 +609,7 @@ class TemplatePatchIn(BaseModel):
display_fields: list[str] | None = None display_fields: list[str] | None = None
row_filters: list[dict[str, Any]] | None = None row_filters: list[dict[str, Any]] | None = None
field_rules: list[FieldRuleIn] | None = None field_rules: list[FieldRuleIn] | None = None
ignore_port_changes: bool | None = None
class MappingRowIn(BaseModel): class MappingRowIn(BaseModel):

View file

@ -0,0 +1,159 @@
"""Template validation, rule semantics, protected references and bounded queries."""
from copy import deepcopy
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 BizCompareJob, BizCompareRun, BizCompareTemplate, BizStateBatchCommand, BizStateBatch
from netx_api.biz_state import compare_service as svc
from netx_api.biz_state.compare_engine import compare_rows, row_key
from netx_api.biz_state.compare_rules import eval_leaf_filter, values_equal
from netx_api.biz_state.compare_sql import compile_row_filters_sql, _rk_sql, _field_differs_sql
from netx_api.biz_state_router import TemplateIn, TemplatePatchIn
@pytest.fixture
def db():
engine = create_engine("sqlite+pysqlite:///:memory:")
Base.metadata.create_all(engine)
session = sessionmaker(bind=engine, expire_on_commit=False)()
yield session
session.close()
engine.dispose()
def body():
return {"name": "Test", "metrics": [{"sheet_id": "s", "metric_id": "custom", "key_fields": ["id"],
"compare_fields": ["count"], "iface_fields": ["port"], "ignore_port_changes": False,
"row_filters": [], "field_rules": []}]}
@pytest.mark.parametrize("bad", [
{"key_fields": []}, {"key_fields": ["id", "id"]}, {"metric_id": ""},
{"row_filters": [{"field": "state", "op": "typo", "value": "up"}]},
{"row_filters": [{"field": "state", "op": "regex", "value": "["}]},
{"row_filters": [{"any": []}]}, {"row_filters": [{"field": "id", "op": "in", "value": "a,b"}]},
{"field_rules": [{"field": "count", "compare": "typo"}]},
{"field_rules": [{"field": "count", "normalize": "typo"}]},
{"field_rules": [{"field": "count", "tolerance": -1}]},
{"field_rules": [{"field": "count", "tolerance": float("inf")}]},
{"field_rules": [{"field": "count", "tolerance": float("nan")}]},
{"field_rules": [{"field": "count"}, {"field": "count"}]},
])
def test_invalid_sheet_never_silently_drops_or_saves(db, bad):
data = body()
data["metrics"].append({**deepcopy(data["metrics"][0]), "sheet_id": "second", **bad})
with pytest.raises(HTTPException) as exc:
svc.create_template(db, data)
assert exc.value.status_code == 400
assert "metrics[1]" in exc.value.detail["path"]
assert db.query(BizCompareTemplate).count() == 0
def test_template_validation_precedes_mutation_and_rejects_duplicate_alias(db):
tpl = svc.create_template(db, body())
data = body()
data["name"] = "Wrong name"
data["metrics"][0]["key_fields"] = []
with pytest.raises(HTTPException):
svc.update_template(db, tpl["id"], data)
assert db.get(BizCompareTemplate, tpl["id"]).name == "Test"
with pytest.raises(HTTPException):
svc.create_template(db, {**body(), "iface_normalize_rules": [{"from": "GE", "to": "a"}, {"from": "ge", "to": "b"}]})
def test_api_port_policy_roundtrip_and_legacy_patch(db):
saved = svc.create_template(db, TemplateIn.model_validate(body()).model_dump())
assert saved["metrics"][0]["ignore_port_changes"] is False
updated = svc.update_template(db, saved["id"], TemplatePatchIn(ignore_port_changes=True).model_dump(exclude_unset=True))
assert updated["metrics"][0]["ignore_port_changes"] is True
legacy = svc.create_template(db, {"name": "legacy", "metric_id": "custom", "key_fields": ["id"], "ignore_port_changes": False})
assert legacy["metrics"][0]["ignore_port_changes"] is False
@pytest.mark.parametrize("value", [0, False])
def test_filter_and_identity_preserve_false_and_zero(value):
assert not eval_leaf_filter({"x": value}, {"field": "x", "op": "empty"})
assert eval_leaf_filter({"x": value}, {"field": "x", "op": "eq", "value": value})
assert row_key({"id": value}, ["id"]) != row_key({"id": None}, ["id"])
result = compare_rows(before_rows=[{"id": value}], after_rows=[{"id": None}], key_fields=["id"], iface_fields=[], compare_fields=[], port_map={})
assert result["summary"]["removed"] == 1 and result["summary"]["added"] == 1
_, params = compile_row_filters_sql([{"field": "x", "op": "eq", "value": value}])
assert list(params.values()) == ["0" if value is not False else "false"]
@pytest.mark.parametrize("mode", ["numeric", "percent"])
def test_missing_numeric_value_is_not_zero(mode):
assert not values_equal(None, 0, rule={"compare": mode, "tolerance": 5})
assert values_equal(None, "", rule={"compare": mode})
assert values_equal("+.5", "0.5", rule={"compare": mode})
assert values_equal("1.", "1", rule={"compare": mode})
assert values_equal("NaN", "NaN", rule={"compare": mode}) # textual fallback
def test_sql_filter_literal_contains_and_collision_free_composite_keys():
sql, params = compile_row_filters_sql([{"field": "x", "op": "contains", "value": "a%_\\b"}])
assert "ESCAPE" in sql
assert list(params.values()) == ["%a\\%\\_\\\\b%"]
sql, _ = compile_row_filters_sql([{"field": "x", "op": "contains", "value": ""}])
assert "FALSE" in sql
assert "jsonb_build_array" in _rk_sql(["a", "b"])
assert "concat_ws" not in _rk_sql(["a", "b"])
assert "[+-]?" in _field_differs_sql("b", "a", {"compare": "numeric"})
def test_sheet_order_uses_one_aggregate_query_without_raw_text(db):
for bid, metric, count in (("b", "large", 20), ("a", "large", 30), ("b", "small", 3)):
db.add(BizStateBatchCommand(id=f"{bid}-{metric}", batch_id=bid, metric_id=metric, row_count=count, raw_text="x" * 100000))
db.commit()
sql = []
event.listen(db.get_bind(), "before_cursor_execute", lambda c, cur, statement, p, ctx, many: sql.append(statement))
sheets = [{"metric_id": "large", "sheet_id": "l1"}, {"metric_id": "small", "sheet_id": "s"}, {"metric_id": "large", "sheet_id": "l2"}]
ordered = svc._order_sheets_small_first(db, sheets, before_batch_id="b", after_batch_id="a")
assert [s["sheet_id"] for s in ordered] == ["s", "l1", "l2"]
assert len(sql) == 1 and "raw_text" not in sql[0]
def test_referenced_template_missing_sheets_and_running_delete_protected(db):
tpl = svc.create_template(db, body())
db.add(BizCompareJob(id="j", template_id=tpl["id"], enabled_sheet_ids=["s"]))
db.add(BizCompareRun(id="r", job_id="j", status="running"))
db.commit()
for action in (lambda: svc.delete_template(db, tpl["id"]), lambda: svc.delete_job(db, "j"), lambda: svc.delete_run(db, "r")):
with pytest.raises(HTTPException) as exc:
action()
assert exc.value.status_code == 409
with pytest.raises(HTTPException) as exc:
svc.update_job(db, "j", {"name": "invalid", "enabled_sheet_ids": ["missing"]})
assert exc.value.detail["error"] == "unknown_enabled_sheets"
assert db.get(BizCompareJob, "j").name != "invalid"
def test_template_usage_counts_are_batched(db, monkeypatch):
monkeypatch.setattr(svc, "ensure_default_templates", lambda db: None)
monkeypatch.setattr(svc, "upgrade_builtin_split_sheets", lambda db: None)
used = svc.create_template(db, body())
unused = svc.create_template(db, {**body(), "name": "Unused"})
db.add_all([BizCompareJob(id="j1", template_id=used["id"]), BizCompareJob(id="j2", template_id=used["id"])])
db.commit()
sql = []
event.listen(db.get_bind(), "before_cursor_execute", lambda c, cur, statement, p, ctx, many: sql.append(statement))
counts = {t["id"]: t["job_count"] for t in svc.list_templates(db)}
assert counts == {used["id"]: 2, unused["id"]: 0}
assert len(sql) == 2 # one template query and one grouped reference query
@pytest.mark.parametrize("status", ["queued", "running", "failed", "cancelled", "partial"])
def test_compare_rejects_incomplete_sources_before_enqueue(db, status):
tpl = svc.create_template(db, body())
db.add(BizStateBatch(id="before", status="success"))
db.add(BizStateBatch(id="after", status=status))
db.add(BizCompareJob(id="j", template_id=tpl["id"], before_batch_id="before", after_batch_id="after"))
db.commit()
with pytest.raises(HTTPException) as exc:
svc._validate_compare_job(db, "j")
assert exc.value.detail == "source_batch_not_complete"
assert db.query(BizCompareRun).count() == 0