From 8d3c443159335f869c5b3b617b045413ac46d53e Mon Sep 17 00:00:00 2001 From: Christian Manivong Date: Mon, 20 Jul 2026 15:05:53 +0200 Subject: [PATCH] feat(firewall): implement apply_firewall_rule + commit_firewall_rules OPNsense-specific half of the FirewallDriver diff/apply mechanism added in napalm-device-types: translates the vendor-neutral rule dict into the /api/firewall/filter/addRule or setRule/ payload (string "1"/"0" booleans, empty interface = floating rule -- same shape as the existing SNMP self-provisioning rule in _action_fix_snmp), and commit_firewall_rules reloads the filter via /api/firewall/filter/apply. get_firewall_rules() already returns compatible field names, no changes needed there. --- napalm_opnsense/opnsense.py | 43 +++++++++++++++++++ tests/unit/test_driver.py | 86 +++++++++++++++++++++++++++++++++++++ 2 files changed, 129 insertions(+) diff --git a/napalm_opnsense/opnsense.py b/napalm_opnsense/opnsense.py index edc5846..6350f1a 100644 --- a/napalm_opnsense/opnsense.py +++ b/napalm_opnsense/opnsense.py @@ -2149,6 +2149,49 @@ class OPNsenseDriver(FirewallDriver): return sorted(result, key=lambda x: (x["floating"], x["is_group"], x["interface"], x["sequence"])) + def apply_firewall_rule(self, rule: dict[str, Any], *, uuid: str | None = None) -> dict[str, Any]: + """Create or update a single OPNsense firewall filter rule. + + `rule` uses the vendor-neutral field names from + ``napalm_device_types.models.FirewallRuleDict`` (see + ``FirewallDriver.diff_firewall_rules``/``apply_firewall_ruleset``, + the generic reconciliation engine that calls this method). This is + the OPNsense-specific half: translating those fields into the + ``/api/firewall/filter/addRule``/``setRule`` payload shape (string + "1"/"0" booleans, empty ``interface`` means a floating rule — same + payload shape as the SNMP self-provisioning rule in + ``_action_fix_snmp``). + """ + payload = { + "rule": { + "enabled": "1" if rule.get("enabled", True) else "0", + "sequence": "1", + "action": rule.get("action", "pass"), + "quick": "1" if rule.get("quick", True) else "0", + "interface": rule.get("interface", "") or "", + "direction": rule.get("direction", "in"), + "ipprotocol": "inet", + "protocol": rule.get("protocol", "any"), + "source_net": rule.get("source_net", "any") or "any", + "source_port": rule.get("source_port", "") or "", + "destination_net": rule.get("destination_net", "any") or "any", + "destination_port": rule.get("destination_port", "") or "", + "log": "1" if rule.get("log", False) else "0", + "floating": "yes" if not rule.get("interface") else "no", + "descr": rule.get("description", ""), + } + } + path = f"/api/firewall/filter/setRule/{uuid}" if uuid else "/api/firewall/filter/addRule" + return self._post(path, payload) + + def commit_firewall_rules(self) -> dict[str, Any]: + """Reload the firewall filter to activate pending rule changes. + + Final step after one or more `apply_firewall_rule()` calls — same + as the last step of `_action_fix_snmp`. + """ + return self._post("/api/firewall/filter/apply", {}) + # ------------------------------------------------------------------ # Hostname management # ------------------------------------------------------------------ diff --git a/tests/unit/test_driver.py b/tests/unit/test_driver.py index 0c8f9cb..5d8a6b0 100644 --- a/tests/unit/test_driver.py +++ b/tests/unit/test_driver.py @@ -1830,3 +1830,89 @@ class TestDeleteDhcpReservationAndLease: with pytest.raises(RuntimeError, match="Kea DHCPv4 plugin unavailable"): driver.delete_dhcp_reservation_and_lease(mac="02:aa:bb:cc:dd:ee", ip="172.22.8.253") + + +# --------------------------------------------------------------------------- +# apply_firewall_rule / commit_firewall_rules +# --------------------------------------------------------------------------- + +class TestApplyFirewallRule: + def test_add_calls_addRule_with_translated_fields(self, driver): + calls = [] + driver._post = lambda path, data=None: calls.append((path, data)) or {"result": "saved"} + + rule = { + "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, + } + driver.apply_firewall_rule(rule) + + assert len(calls) == 1 + path, payload = calls[0] + assert path == "/api/firewall/filter/addRule" + assert payload["rule"]["enabled"] == "1" + assert payload["rule"]["quick"] == "1" + assert payload["rule"]["log"] == "0" + assert payload["rule"]["action"] == "pass" + assert payload["rule"]["interface"] == "lan" + assert payload["rule"]["source_net"] == "MGMT_NET" + assert payload["rule"]["destination_net"] == "(self)" + assert payload["rule"]["destination_port"] == "https" + assert payload["rule"]["descr"] == "allow_mgmt_to_fw_gui" + assert payload["rule"]["floating"] == "no" + + def test_update_calls_setRule_with_uuid(self, driver): + calls = [] + driver._post = lambda path, data=None: calls.append((path, data)) or {"result": "saved"} + + driver.apply_firewall_rule({"description": "x", "action": "pass"}, uuid="abc-123") + + assert calls[0][0] == "/api/firewall/filter/setRule/abc-123" + + def test_empty_interface_is_a_floating_rule(self, driver): + calls = [] + driver._post = lambda path, data=None: calls.append((path, data)) or {"result": "saved"} + + driver.apply_firewall_rule({"description": "x", "action": "pass", "interface": ""}) + + assert calls[0][1]["rule"]["floating"] == "yes" + + def test_disabled_quick_log_flags_translate_to_zero(self, driver): + calls = [] + driver._post = lambda path, data=None: calls.append((path, data)) or {"result": "saved"} + + driver.apply_firewall_rule( + { + "description": "x", + "action": "block", + "enabled": False, + "quick": False, + "log": True, + } + ) + + rule = calls[0][1]["rule"] + assert rule["enabled"] == "0" + assert rule["quick"] == "0" + assert rule["log"] == "1" + + +class TestCommitFirewallRules: + def test_calls_filter_apply(self, driver): + calls = [] + driver._post = lambda path, data=None: calls.append((path, data)) or {"status": "ok"} + + result = driver.commit_firewall_rules() + + assert calls == [("/api/firewall/filter/apply", {})] + assert result == {"status": "ok"}