mirror of
https://github.com/hansjone/netx.git
synced 2026-10-12 04:10:47 +08:00
fix(biz-compare): validate templates and protect comparison inputs
This commit is contained in:
parent
8c15975d3e
commit
90cf1907b4
7 changed files with 376 additions and 54 deletions
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
130
netx_api/biz_state/compare_validation.py
Normal file
130
netx_api/biz_state/compare_validation.py
Normal 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)
|
||||||
|
|
@ -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):
|
||||||
|
|
|
||||||
159
tests/test_biz_compare_templates.py
Normal file
159
tests/test_biz_compare_templates.py
Normal 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
|
||||||
Loading…
Add table
Add a link
Reference in a new issue