From 90cf1907b426806d05ac6adff37f4e3fae5e3ef8 Mon Sep 17 00:00:00 2001 From: oliver Date: Sun, 11 Oct 2026 05:56:09 +0800 Subject: [PATCH] fix(biz-compare): validate templates and protect comparison inputs --- netx_api/biz_state/compare_engine.py | 4 +- netx_api/biz_state/compare_rules.py | 32 +++-- netx_api/biz_state/compare_service.py | 78 ++++++----- netx_api/biz_state/compare_sql.py | 24 ++-- netx_api/biz_state/compare_validation.py | 130 ++++++++++++++++++ netx_api/biz_state_router.py | 3 + tests/test_biz_compare_templates.py | 159 +++++++++++++++++++++++ 7 files changed, 376 insertions(+), 54 deletions(-) create mode 100644 netx_api/biz_state/compare_validation.py create mode 100644 tests/test_biz_compare_templates.py diff --git a/netx_api/biz_state/compare_engine.py b/netx_api/biz_state/compare_engine.py index 5957d79..87af097 100644 --- a/netx_api/biz_state/compare_engine.py +++ b/netx_api/biz_state/compare_engine.py @@ -7,7 +7,7 @@ from heapq import heappush, heapreplace from itertools import chain 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 ( apply_iface_normalize_rows, normalize_iface_rules, @@ -94,7 +94,7 @@ def apply_port_map( 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( diff --git a/netx_api/biz_state/compare_rules.py b/netx_api/biz_state/compare_rules.py index 97e454b..269a32a 100644 --- a/netx_api/biz_state/compare_rules.py +++ b/netx_api/biz_state/compare_rules.py @@ -7,10 +7,21 @@ All metric-specific compare behavior belongs in the sheet template from __future__ import annotations import re +import math from typing import Any, Mapping, Sequence _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]: @@ -22,7 +33,7 @@ def _as_list(raw: Any) -> list[Any]: 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: @@ -35,17 +46,17 @@ def eval_leaf_filter(row: Mapping[str, Any], filt: Mapping[str, Any]) -> bool: expect = filt.get("value") if op in ("eq", "=="): - return raw.lower() == str(expect or "").strip().lower() + return raw.lower() == scalar_text(expect).strip().lower() if op in ("ne", "!="): - return raw.lower() != str(expect or "").strip().lower() + return raw.lower() != scalar_text(expect).strip().lower() 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 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 if op == "contains": - needle = str(expect or "").strip().lower() + needle = scalar_text(expect).strip().lower() return bool(needle) and needle in raw.lower() if op == "empty": 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 return bool(_AGE_TIMER_RE.match(raw)) if op == "ci_eq": - return raw.lower() == str(expect or "").strip().lower() + return raw.lower() == scalar_text(expect).strip().lower() return True @@ -99,7 +110,7 @@ def apply_row_filters( 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() if not mode or mode in ("none", "strip"): return text @@ -190,8 +201,11 @@ def effective_display_fields( def _parse_float(text: str) -> float | None: + if not re.fullmatch(NUMBER_PATTERN, text): + return None try: - return float(text) if text else 0.0 + value = float(text) + return value if math.isfinite(value) else None except ValueError: return None diff --git a/netx_api/biz_state/compare_service.py b/netx_api/biz_state/compare_service.py index 43e097b..4ac8ddd 100644 --- a/netx_api/biz_state/compare_service.py +++ b/netx_api/biz_state/compare_service.py @@ -45,6 +45,7 @@ from .iface_normalize import ( normalize_iface_rules, ) from .profiles import metric_field_map +from .compare_validation import validate_template_body _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": return True cmds = ( - db.query(BizStateBatchCommand) + db.query(BizStateBatchCommand.parse_status, BizStateBatchCommand.row_count) .filter( BizStateBatchCommand.batch_id == bid, 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( db: Session, 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.""" if len(sheets) <= 1: 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]]] = [] 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.sort(key=lambda x: (x[0], x[1])) 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, row_filters=_normalize_row_filters(body.get("row_filters")), 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) upgrade_builtin_split_sheets(db) 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]]: @@ -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]: + validate_template_body(body) sheets = _parse_metrics_body(body) t = BizCompareTemplate( 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) if not t: raise HTTPException(status_code=404, detail="template_not_found") + validate_template_body(body, partial=True) if "name" in body: t.name = str(body.get("name") or "")[:256] if "note" in body: @@ -1485,6 +1471,7 @@ def update_template(db: Session, template_id: str, body: dict[str, Any]) -> dict "display_fields", "row_filters", "field_rules", + "ignore_port_changes", ) ): # 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")) if "field_rules" in body: 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 _apply_sheets_to_row(t, sheets) t.updated_at = _utcnow() @@ -1533,6 +1522,8 @@ def delete_template(db: Session, template_id: str) -> None: t = db.get(BizCompareTemplate, template_id) if not t: 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.commit() @@ -1888,6 +1879,16 @@ def list_jobs(db: Session) -> list[dict[str, Any]]: 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]: ensure_default_templates(db) 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): raise HTTPException(status_code=404, detail="template_not_found") enabled = _str_list(body.get("enabled_sheet_ids")) + _validate_job_sheets(db, template_id, enabled) j = BizCompareJob( id=uuid4().hex, 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) if not j: 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 ( "name", "template_id", @@ -1953,6 +1957,9 @@ def delete_job(db: Session, job_id: str) -> None: j = db.get(BizCompareJob, job_id) if not j: 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 = [ 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) if not r: 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 "") db.query(BizCompareDiff).filter(BizCompareDiff.run_id == run_id).delete( 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) if not before_batch_id or not after_batch_id: 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") + 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) if not sheets_cfg: diff --git a/netx_api/biz_state/compare_sql.py b/netx_api/biz_state/compare_sql.py index b7ba9d3..672a15a 100644 --- a/netx_api/biz_state/compare_sql.py +++ b/netx_api/biz_state/compare_sql.py @@ -14,7 +14,7 @@ from uuid import uuid4 from sqlalchemy import text 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") @@ -53,7 +53,7 @@ _SQL_COMPARE_MODES = frozenset( } ) _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: @@ -258,27 +258,30 @@ def compile_row_filters_sql( expect = filt.get("value") if op in ("eq", "==", "ci_eq"): key = _next("v") - params[key] = str(expect or "").strip() + params[key] = scalar_text(expect).strip() if op == "ci_eq" or op in ("eq", "=="): # 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})" if op in ("ne", "!="): key = _next("v") - params[key] = str(expect or "").strip().lower() + params[key] = scalar_text(expect).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})" + needle = scalar_text(expect).strip().lower() + if not needle: + 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"): 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()] + clean = [scalar_text(x).strip().lower() for x in vals if scalar_text(x).strip()] if not clean: return "FALSE" if op == "in" else "TRUE" 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}', ''))") if len(parts) == 1: 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: diff --git a/netx_api/biz_state/compare_validation.py b/netx_api/biz_state/compare_validation.py new file mode 100644 index 0000000..442b10f --- /dev/null +++ b/netx_api/biz_state/compare_validation.py @@ -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) diff --git a/netx_api/biz_state_router.py b/netx_api/biz_state_router.py index 9c59382..fa1e936 100644 --- a/netx_api/biz_state_router.py +++ b/netx_api/biz_state_router.py @@ -574,6 +574,7 @@ class TemplateMetricIn(BaseModel): display_fields: list[str] = Field(default_factory=list) row_filters: list[dict[str, Any]] = Field(default_factory=list) field_rules: list[FieldRuleIn] = Field(default_factory=list) + ignore_port_changes: bool | None = None class TemplateIn(BaseModel): @@ -592,6 +593,7 @@ class TemplateIn(BaseModel): display_fields: list[str] = Field(default_factory=list) row_filters: list[dict[str, Any]] = Field(default_factory=list) field_rules: list[FieldRuleIn] = Field(default_factory=list) + ignore_port_changes: bool | None = None class TemplatePatchIn(BaseModel): @@ -607,6 +609,7 @@ class TemplatePatchIn(BaseModel): display_fields: list[str] | None = None row_filters: list[dict[str, Any]] | None = None field_rules: list[FieldRuleIn] | None = None + ignore_port_changes: bool | None = None class MappingRowIn(BaseModel): diff --git a/tests/test_biz_compare_templates.py b/tests/test_biz_compare_templates.py new file mode 100644 index 0000000..2db9109 --- /dev/null +++ b/tests/test_biz_compare_templates.py @@ -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