From 90b8e087890f1b790d51fbbc858ce4bfcc774640 Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Mon, 20 Jul 2026 15:01:13 +0200 Subject: [PATCH] feat(firewall): add generic diff/apply mechanism for firewall rules FirewallRuleDict/FirewallRuleDiffDict (models.py) plus three abstract methods (get_firewall_rules/apply_firewall_rule/commit_firewall_rules) concrete drivers implement, and two concrete methods every driver gets for free: diff_firewall_rules() matches desired vs. live rules by description and reports add/update (never delete -- a firewall may carry manually-created rules a caller's desired set was never meant to describe); apply_firewall_ruleset() orchestrates applying the diff and yields progress lines, meant for streaming to a caller. This is the generic reconciliation engine NetOrk's Firewall Profile feature needs against OPNsense -- kept here instead of in napalm-opnsense since the matching/comparison/orchestration logic is identical for any firewall vendor that implements the three abstract methods. --- napalm_device_types/firewall.py | 147 +++++++++++++++++++++++++- napalm_device_types/models.py | 39 +++++++ tests/test_firewall_diff_apply.py | 170 ++++++++++++++++++++++++++++++ 3 files changed, 355 insertions(+), 1 deletion(-) create mode 100644 tests/test_firewall_diff_apply.py diff --git a/napalm_device_types/firewall.py b/napalm_device_types/firewall.py index 9507e25..8fee538 100644 --- a/napalm_device_types/firewall.py +++ b/napalm_device_types/firewall.py @@ -10,10 +10,13 @@ Usage:: ... """ -from typing import Any, Dict, List +from typing import Any, Dict, Iterator, List, Optional from napalm_device_types.base import DeviceTypeDriver from napalm_device_types._ucd_metrics import IF_SKIP_DEFAULT, collect_ucd_metrics from napalm_device_types.models import ( + FirewallRuleDict, + FirewallRuleDiffDict, + FirewallRuleUpdateDict, HealthMetricsDict, NATTranslationDict, PackageDict, @@ -22,6 +25,20 @@ from napalm_device_types.models import ( VPNTunnelDict, ) +_FIREWALL_RULE_COMPARE_FIELDS = ( + "action", + "interface", + "direction", + "protocol", + "source_net", + "source_port", + "destination_net", + "destination_port", + "log", + "quick", + "enabled", +) + class FirewallDriver(DeviceTypeDriver): TYPE_LABEL: str = "Firewall" @@ -325,3 +342,131 @@ class FirewallDriver(DeviceTypeDriver): ) """ raise NotImplementedError + + # ------------------------------------------------------------------ + # Firewall rule diff/apply. get_firewall_rules/apply_firewall_rule/ + # commit_firewall_rules are abstract (device communication); everything + # else here is a concrete, vendor-neutral algorithm -- see README.md + # "Design principle: generic vs. device-specific logic". + # ------------------------------------------------------------------ + + def get_firewall_rules(self) -> List[FirewallRuleDict]: + """ + Returns all firewall filter rules currently configured on the device. + + `description` must be a stable, human-assigned identifier -- it is + the key used to match rules across calls (most firewall vendors + don't expose an ID a caller can pre-assign). + + :raises NotImplementedError: If the driver does not support reading + firewall rules. + """ + raise NotImplementedError + + def apply_firewall_rule( + self, rule: FirewallRuleDict, *, uuid: Optional[str] = None + ) -> Dict[str, Any]: + """ + Creates or updates a single firewall filter rule on the device. + + :param rule: The desired rule state, in vendor-neutral form. + :param uuid: If given, update the existing rule with this ID + in-place. If ``None``, create a new rule. + :raises NotImplementedError: If the driver does not support writing + firewall rules. + :raises ValueError: If `rule` references an alias/interface the + device doesn't know about. + :raises RuntimeError: If the device rejects the write. + + :returns: A dict with at least ``{"success": bool}``. + """ + raise NotImplementedError + + def commit_firewall_rules(self) -> Dict[str, Any]: + """ + Applies pending firewall filter rule changes (e.g. reloads pf/pfctl, + or whatever the device's equivalent of "Apply Changes" is). + + Call once after one or more `apply_firewall_rule()` calls -- not + after every single rule. + + :raises NotImplementedError: If the driver does not support this + (e.g. rules take effect immediately on write). + + :returns: A dict with at least ``{"success": bool}``. + """ + raise NotImplementedError + + def diff_firewall_rules(self, desired: List[FirewallRuleDict]) -> FirewallRuleDiffDict: + """ + Compares `desired` against the device's current rules and returns + what would need to change to reach that state. + + Matches rules by `description`. A desired rule with no live + counterpart becomes an "add"; a live rule whose description matches + but whose other fields differ becomes an "update". Live rules with + no matching desired entry are **not** reported for deletion -- this + is intentionally conservative: a firewall may carry manually-created + or otherwise unmanaged rules that a caller's `desired` set was never + meant to describe, and this method has no way to distinguish those + from ones simply no longer wanted. Callers wanting delete/cleanup + semantics must implement that themselves, deliberately. + + :param desired: The complete desired rule set. + :returns: ``{"add": [...], "update": [{"uuid", "rule", + "changed_fields"}, ...]}``. + """ + live_by_description: Dict[str, FirewallRuleDict] = { + rule["description"]: rule for rule in self.get_firewall_rules() + } + + add: List[FirewallRuleDict] = [] + update: List[FirewallRuleUpdateDict] = [] + + for desired_rule in desired: + live_rule = live_by_description.get(desired_rule["description"]) + if live_rule is None: + add.append(desired_rule) + continue + + changed_fields = [ + field + for field in _FIREWALL_RULE_COMPARE_FIELDS + if live_rule.get(field) != desired_rule.get(field) + ] + if changed_fields: + update.append( + { + "uuid": live_rule["uuid"], + "rule": desired_rule, + "changed_fields": changed_fields, + } + ) + + return {"add": add, "update": update} + + def apply_firewall_ruleset(self, desired: List[FirewallRuleDict]) -> Iterator[str]: + """ + Computes the diff against `desired` and applies it, yielding one + human-readable progress line per change, then commits. + + Intended for streaming to a caller (e.g. an SSE endpoint) that wants + live progress while writing to a real device. + + :param desired: The complete desired rule set. + :yields: Progress lines, one per applied add/update, plus a final + commit line. + """ + diff = self.diff_firewall_rules(desired) + + for rule in diff["add"]: + self.apply_firewall_rule(rule) + yield f"[add] {rule['description']}" + + for entry in diff["update"]: + self.apply_firewall_rule(entry["rule"], uuid=entry["uuid"]) + fields = ", ".join(entry["changed_fields"]) + yield f"[update] {entry['rule']['description']} ({fields})" + + self.commit_firewall_rules() + yield f"[commit] applied {len(diff['add'])} add(s), {len(diff['update'])} update(s)" diff --git a/napalm_device_types/models.py b/napalm_device_types/models.py index 4ff56ca..6097a38 100644 --- a/napalm_device_types/models.py +++ b/napalm_device_types/models.py @@ -314,6 +314,45 @@ class SessionDict(TypedDict): age: float +class FirewallRuleDict(TypedDict): + """A single firewall filter rule, in vendor-neutral form. + + `description` is the stable matching key across get_firewall_rules()/ + diff_firewall_rules()/apply_firewall_rule() -- firewall vendors + generally don't expose an ID a caller can pre-assign, so the rule's + human description is what ties a "desired" rule to its "live" + counterpart. `source_net`/`source_port`/`destination_net`/ + `destination_port` are plain strings (comma-joined by the caller if a + rule references multiple aliases) -- driver methods never expand or + split them. + """ + + uuid: str + description: str + action: str + interface: str + direction: str + protocol: str + source_net: str + source_port: str + destination_net: str + destination_port: str + enabled: bool + quick: bool + log: bool + + +class FirewallRuleUpdateDict(TypedDict): + uuid: str + rule: FirewallRuleDict + changed_fields: List[str] + + +class FirewallRuleDiffDict(TypedDict): + add: List[FirewallRuleDict] + update: List[FirewallRuleUpdateDict] + + class VPNTunnelDict(TypedDict): type: str local_endpoint: str diff --git a/tests/test_firewall_diff_apply.py b/tests/test_firewall_diff_apply.py new file mode 100644 index 0000000..b65ccf6 --- /dev/null +++ b/tests/test_firewall_diff_apply.py @@ -0,0 +1,170 @@ +"""Tests for FirewallDriver's generic diff/apply mechanism. + +diff_firewall_rules/apply_firewall_ruleset are concrete methods on the base +class (not overridden by concrete drivers) — they only depend on the three +abstract methods (get_firewall_rules/apply_firewall_rule/commit_firewall_rules), +so a fake in-memory driver is enough to exercise them fully; no real device +or vendor driver needed. See README.md "Design principle: generic vs. +device-specific logic" for why this logic lives here and not in a vendor +driver. +""" + +from typing import Any, Dict, List, Optional + +import pytest +from napalm_device_types import FirewallDriver +from napalm_device_types.models import FirewallRuleDict + + +def _rule(**overrides: Any) -> FirewallRuleDict: + base: FirewallRuleDict = { + "uuid": "", + "description": "allow_mgmt_to_fw_gui", + "action": "pass", + "interface": "lan", + "direction": "in", + "protocol": "tcp", + "source_net": "MGMT_NET", + "source_port": "", + "destination_net": "(self)", + "destination_port": "https", + "enabled": True, + "quick": True, + "log": False, + } + base.update(overrides) # type: ignore[typeddict-item] + return base + + +class _FakeFirewall(FirewallDriver): + """In-memory fake — no network, no OPNsense/vendor specifics.""" + + def __init__(self, live_rules: Optional[List[FirewallRuleDict]] = None) -> None: + self.live_rules = live_rules or [] + self.applied: List[Dict[str, Any]] = [] + self.committed = False + + def get_firewall_rules(self) -> List[FirewallRuleDict]: + return self.live_rules + + def apply_firewall_rule( + self, rule: FirewallRuleDict, *, uuid: Optional[str] = None + ) -> Dict[str, Any]: + self.applied.append({"rule": rule, "uuid": uuid}) + return {"success": True} + + def commit_firewall_rules(self) -> Dict[str, Any]: + self.committed = True + return {"success": True} + + +class TestAbstractContract: + def test_get_firewall_rules_raises_not_implemented_by_default(self): + class _Bare(FirewallDriver): + def __init__(self) -> None: + pass + + with pytest.raises(NotImplementedError): + _Bare().get_firewall_rules() + + def test_apply_firewall_rule_raises_not_implemented_by_default(self): + class _Bare(FirewallDriver): + def __init__(self) -> None: + pass + + with pytest.raises(NotImplementedError): + _Bare().apply_firewall_rule(_rule()) + + def test_commit_firewall_rules_raises_not_implemented_by_default(self): + class _Bare(FirewallDriver): + def __init__(self) -> None: + pass + + with pytest.raises(NotImplementedError): + _Bare().commit_firewall_rules() + + +class TestDiffFirewallRules: + def test_desired_rule_missing_live_is_an_add(self): + driver = _FakeFirewall(live_rules=[]) + diff = driver.diff_firewall_rules([_rule()]) + + assert len(diff["add"]) == 1 + assert diff["add"][0]["description"] == "allow_mgmt_to_fw_gui" + assert diff["update"] == [] + + def test_matching_rule_with_changed_field_is_an_update(self): + live = _rule(uuid="abc-123", action="block") + driver = _FakeFirewall(live_rules=[live]) + diff = driver.diff_firewall_rules([_rule(action="pass")]) + + assert diff["add"] == [] + assert len(diff["update"]) == 1 + update = diff["update"][0] + assert update["uuid"] == "abc-123" + assert update["changed_fields"] == ["action"] + assert update["rule"]["action"] == "pass" + + def test_identical_rule_produces_no_diff(self): + live = _rule(uuid="abc-123") + driver = _FakeFirewall(live_rules=[live]) + diff = driver.diff_firewall_rules([_rule()]) + + assert diff["add"] == [] + assert diff["update"] == [] + + def test_live_rule_not_in_desired_is_never_deleted(self): + """v1 never deletes -- rules present live but absent from `desired` + are simply ignored, not reported for removal.""" + live = _rule(uuid="abc-123", description="some_unmanaged_rule") + driver = _FakeFirewall(live_rules=[live]) + diff = driver.diff_firewall_rules([]) + + assert diff == {"add": [], "update": []} + assert "delete" not in diff + + def test_multiple_changed_fields_all_reported(self): + live = _rule(uuid="abc-123", action="block", protocol="udp", log=True) + driver = _FakeFirewall(live_rules=[live]) + diff = driver.diff_firewall_rules([_rule(action="pass", protocol="tcp", log=False)]) + + assert set(diff["update"][0]["changed_fields"]) == {"action", "protocol", "log"} + + +class TestApplyFirewallRuleset: + def test_adds_are_applied_with_no_uuid(self): + driver = _FakeFirewall(live_rules=[]) + list(driver.apply_firewall_ruleset([_rule()])) + + assert len(driver.applied) == 1 + assert driver.applied[0]["uuid"] is None + assert driver.applied[0]["rule"]["description"] == "allow_mgmt_to_fw_gui" + + def test_updates_are_applied_with_existing_uuid(self): + live = _rule(uuid="abc-123", action="block") + driver = _FakeFirewall(live_rules=[live]) + list(driver.apply_firewall_ruleset([_rule(action="pass")])) + + assert len(driver.applied) == 1 + assert driver.applied[0]["uuid"] == "abc-123" + + def test_commits_after_applying(self): + driver = _FakeFirewall(live_rules=[]) + list(driver.apply_firewall_ruleset([_rule()])) + + assert driver.committed is True + + def test_yields_a_progress_line_per_change(self): + driver = _FakeFirewall(live_rules=[]) + lines = list(driver.apply_firewall_ruleset([_rule(), _rule(description="second_rule")])) + + assert len(lines) >= 2 + assert all(isinstance(line, str) for line in lines) + + def test_no_changes_still_commits_but_applies_nothing(self): + live = _rule(uuid="abc-123") + driver = _FakeFirewall(live_rules=[live]) + list(driver.apply_firewall_ruleset([_rule()])) + + assert driver.applied == [] + assert driver.committed is True