mirror of
https://github.com/hansjone/netx.git
synced 2026-10-11 23:20:52 +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
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