Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Empty file.
9 changes: 9 additions & 0 deletions custom_components/pfsense/config_flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,11 @@
CONF_DEVICE_TRACKER_ENABLED,
CONF_DEVICE_TRACKER_SCAN_INTERVAL,
CONF_DEVICES,
CONF_RULE_SWITCH_KILL_STATES,
DEFAULT_DEVICE_TRACKER_CONSIDER_HOME,
DEFAULT_DEVICE_TRACKER_ENABLED,
DEFAULT_DEVICE_TRACKER_SCAN_INTERVAL,
DEFAULT_RULE_SWITCH_KILL_STATES,
DEFAULT_SCAN_INTERVAL,
DEFAULT_VERIFY_SSL,
DOMAIN,
Expand Down Expand Up @@ -235,6 +237,13 @@ async def async_step_init(self, user_input=None):
DEFAULT_DEVICE_TRACKER_CONSIDER_HOME,
),
): vol.All(vol.Coerce(int), vol.Clamp(min=0, max=600)),
vol.Optional(
CONF_RULE_SWITCH_KILL_STATES,
default=opts.get(
CONF_RULE_SWITCH_KILL_STATES,
DEFAULT_RULE_SWITCH_KILL_STATES,
),
): bool,
}
return self.async_show_form(step_id="init", data_schema=vol.Schema(base_schema))

Expand Down
3 changes: 3 additions & 0 deletions custom_components/pfsense/const.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,9 @@

CONF_DEVICES = "devices"

CONF_RULE_SWITCH_KILL_STATES = "rule_switch_kill_states"
DEFAULT_RULE_SWITCH_KILL_STATES = False

COUNT = "count"

BYTES_RECEIVED = "bytes_received"
Expand Down
118 changes: 115 additions & 3 deletions custom_components/pfsense/pypfsense/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from __future__ import annotations

import asyncio
import ipaddress
import logging
import re
from typing import Any
Expand Down Expand Up @@ -343,6 +344,16 @@ async def get_nat_outbound_rules(self) -> list[dict]:
data = await self._get("/firewall/nat/outbound/mappings")
return data or []

# pfSense's own webConfigurator stores content-less config elements
# (``<statetype></statetype>``) that the REST API surfaces as empty strings.
# A PATCH re-validates the whole rule object, so those empty-but-required
# fields make an otherwise unrelated toggle fail with
# ``FIELD_EMPTY_NOT_ALLOWED``. Re-send them with the value pfSense would
# have defaulted to, which is a no-op for the rule's behaviour.
_RULE_REQUIRED_DEFAULTS = {
"/firewall/rule": {"statetype": "keep state"},
}

async def _set_rule_disabled(
self, path: str, rules: list[dict], match_key: str, match_value, disabled: bool
) -> None:
Expand All @@ -351,10 +362,12 @@ async def _set_rule_disabled(
continue
if bool(rule.get("disabled")) == disabled:
return
payload = {"id": rule["id"], "disabled": disabled}
for field, fallback in self._RULE_REQUIRED_DEFAULTS.get(path, {}).items():
if not rule.get(field):
payload[field] = fallback
async with self._write_lock:
await self._request(
"PATCH", path, payload={"id": rule["id"], "disabled": disabled}
)
await self._request("PATCH", path, payload=payload)
await self._apply("firewall")
return

Expand Down Expand Up @@ -493,6 +506,60 @@ async def kill_states(self, source: str, destination: str | None = None) -> None
cmd += f" -k {_shq(destination)}"
await self.exec_command(cmd)

async def kill_states_for_rule(self, rule: dict) -> None:
"""Drop state-table entries for the hosts/networks a rule matches on.

Used when a rule switch is toggled so existing connections don't keep
flowing under the old ruleset. REST-only, best effort: this uses the
``DELETE /firewall/states`` prefix filter, so rule endpoints it can't
turn into an IPv4 host / octet-aligned CIDR prefix -- ``any``,
``(self)``, interface macros (``wan:ip``), negated aliases, non
octet-aligned networks (e.g. ``/25``), IPv6 -- are skipped rather than
falling back to ``pfctl``.
"""
prefixes: set[str] = set()
for net in await self._rule_match_networks(rule):
prefix = _states_prefix(net)
if prefix:
prefixes.add(prefix)
if not prefixes:
return

async with self._write_lock:
for prefix in prefixes:
for field in ("source", "destination"):
try:
await self._request(
"DELETE",
"/firewall/states",
params={f"{field}__startswith": prefix, "limit": 0},
)
except PfSenseAPIError:
pass

async def _rule_match_networks(self, rule: dict) -> list:
"""Concrete ``ip_network`` objects for a rule's source + destination."""
out: list = []
aliases: list[dict] | None = None
for side in ("source", "destination"):
value = rule.get(side)
if not isinstance(value, str):
continue
value = value.strip()
if not value or value.startswith("!") or value in ("any", "(self)"):
continue
if ":" in value: # interface address macros, e.g. ``wan:ip``
continue
net = _as_network(value)
if net is not None:
out.append(net)
continue
# Otherwise treat it as an alias name and expand it.
if aliases is None:
aliases = await self._get("/firewall/aliases") or []
out.extend(_expand_alias_networks(aliases, value))
return out

# ------------------------------------------------------- system control

async def system_reboot(self, type: str = "normal") -> None:
Expand Down Expand Up @@ -535,6 +602,51 @@ def _shq(value: str) -> str:
return "'" + str(value).replace("'", "'\\''") + "'"


def _as_network(value: str):
"""Parse ``value`` as an IPv4/IPv6 host or CIDR, or ``None``."""
try:
return ipaddress.ip_network(value, strict=False)
except ValueError:
return None


def _states_prefix(net) -> str | None:
"""``str`` that ``firewall/state`` source/destination values start with for
every address in ``net``, or ``None`` if ``net`` can't be expressed that way.

States render endpoints as ``ip:port`` (IPv4) so a host becomes ``"ip:"``
and an octet-aligned network becomes its leading octets plus a dot.
"""
if net.version != 4:
return None
if net.prefixlen == 32:
return f"{net.network_address}:"
if net.prefixlen in (8, 16, 24):
octets = str(net.network_address).split(".")
return ".".join(octets[: net.prefixlen // 8]) + "."
return None


def _expand_alias_networks(aliases: list[dict], name: str, _depth: int = 3) -> list:
"""Flatten a host/network alias (following nested aliases) to networks."""
if _depth <= 0:
return []
target = next((a for a in aliases if a.get("name") == name), None)
if target is None or target.get("type") not in ("host", "network"):
return []
found: list = []
for entry in target.get("address") or []:
entry = str(entry).strip()
if not entry or entry.startswith("!"):
continue
net = _as_network(entry)
if net is not None:
found.append(net)
else: # a nested alias reference
found.extend(_expand_alias_networks(aliases, entry, _depth - 1))
return found


def _flatten_params(params: dict | None) -> dict | None:
"""aiohttp needs str values; drop ``None`` and stringify the rest."""
if not params:
Expand Down
3 changes: 2 additions & 1 deletion custom_components/pfsense/strings.json
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,8 @@
"scan_interval": "Scan Interval (seconds)",
"device_tracker_enabled": "Enable Device Tracker",
"device_tracker_scan_interval": "Device Tracker Scan Interval (seconds)",
"device_tracker_consider_home": "Device Tracker Consider Home (seconds)"
"device_tracker_consider_home": "Device Tracker Consider Home (seconds)",
"rule_switch_kill_states": "Reset matching states when a rule switch is toggled"
}
},
"device_tracker": {
Expand Down
24 changes: 23 additions & 1 deletion custom_components/pfsense/switch.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,12 @@
from homeassistant.util import slugify

from . import CoordinatorEntityManager, PfSenseEntity
from .const import COORDINATOR, DOMAIN
from .const import (
CONF_RULE_SWITCH_KILL_STATES,
COORDINATOR,
DEFAULT_RULE_SWITCH_KILL_STATES,
DOMAIN,
)

_LOGGER = logging.getLogger(__name__)

Expand Down Expand Up @@ -175,6 +180,19 @@ def is_on(self):
def extra_state_attributes(self):
return None

async def _maybe_kill_rule_states(self, rule):
"""Flush the state table for a rule's hosts after a toggle, if enabled."""
if not rule:
return
if not self.config_entry.options.get(
CONF_RULE_SWITCH_KILL_STATES, DEFAULT_RULE_SWITCH_KILL_STATES
):
return
try:
await self._get_pfsense_client().kill_states_for_rule(rule)
except Exception: # best effort - never fail the toggle over this
_LOGGER.warning("failed to kill states for toggled rule", exc_info=True)


class PfSenseFilterSwitch(PfSenseSwitch):
def _pfsense_get_tracker(self):
Expand Down Expand Up @@ -211,6 +229,7 @@ async def async_turn_on(self, **kwargs):
tracker = self._pfsense_get_tracker()
client = self._get_pfsense_client()
await client.enable_filter_rule_by_tracker(tracker)
await self._maybe_kill_rule_states(rule)
await self.coordinator.async_refresh()

async def async_turn_off(self, **kwargs):
Expand All @@ -221,6 +240,7 @@ async def async_turn_off(self, **kwargs):
tracker = self._pfsense_get_tracker()
client = self._get_pfsense_client()
await client.disable_filter_rule_by_tracker(tracker)
await self._maybe_kill_rule_states(rule)
await self.coordinator.async_refresh()


Expand Down Expand Up @@ -274,6 +294,7 @@ async def async_turn_on(self, **kwargs):
method = client.enable_nat_outbound_rule_by_created_time

await method(tracker)
await self._maybe_kill_rule_states(rule)
await self.coordinator.async_refresh()

async def async_turn_off(self, **kwargs):
Expand All @@ -290,6 +311,7 @@ async def async_turn_off(self, **kwargs):
method = client.disable_nat_outbound_rule_by_created_time

await method(tracker)
await self._maybe_kill_rule_states(rule)
await self.coordinator.async_refresh()


Expand Down
3 changes: 2 additions & 1 deletion custom_components/pfsense/translations/en.json
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,8 @@
"scan_interval": "Scan Interval (seconds)",
"device_tracker_enabled": "Enable Device Tracker",
"device_tracker_scan_interval": "Device Tracker Scan Interval (seconds)",
"device_tracker_consider_home": "Device Tracker Consider Home (seconds)"
"device_tracker_consider_home": "Device Tracker Consider Home (seconds)",
"rule_switch_kill_states": "Reset matching states when a rule switch is toggled"
}
},
"device_tracker": {
Expand Down
61 changes: 56 additions & 5 deletions tests/test_pypfsense.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
"""Unit tests for the async pfSense REST API v2 client."""

import re

import aiohttp
import pytest
from aioresponses import aioresponses
Expand Down Expand Up @@ -110,21 +112,70 @@ async def test_carp_status_reduces_to_bool(client):
assert await client.get_carp_status() is False


def _patch_body(m):
return next(
r for (method, url), reqs in m.requests.items()
for r in reqs if method == "PATCH"
).kwargs["json"]


async def test_disable_filter_rule_patches_then_applies(client):
rules = [
{"id": 4, "tracker": 111, "disabled": False, "descr": "r"},
{"id": 5, "tracker": 222, "disabled": False, "descr": "r2"},
{"id": 5, "tracker": 222, "disabled": False, "descr": "r2",
"statetype": "keep state"},
]
with aioresponses() as m:
m.get(f"{API}/firewall/rules", payload=_envelope(rules))
m.patch(f"{API}/firewall/rule", payload=_envelope({"id": 5, "disabled": True}))
m.post(f"{API}/firewall/apply", payload=_envelope({"applied": True}))
await client.disable_filter_rule_by_tracker(222)
req = next(
r for (method, url), reqs in m.requests.items()
for r in reqs if method == "PATCH"
assert _patch_body(m) == {"id": 5, "disabled": True}


async def test_disable_filter_rule_backfills_empty_statetype(client):
# pfSense's GUI writes ``<statetype></statetype>``; a bare disabled PATCH then
# fails FIELD_EMPTY_NOT_ALLOWED, so the client re-sends the default.
rules = [{"id": 5, "tracker": 222, "disabled": False, "statetype": ""}]
with aioresponses() as m:
m.get(f"{API}/firewall/rules", payload=_envelope(rules))
m.patch(f"{API}/firewall/rule", payload=_envelope({"id": 5, "disabled": True}))
m.post(f"{API}/firewall/apply", payload=_envelope({"applied": True}))
await client.disable_filter_rule_by_tracker(222)
assert _patch_body(m) == {
"id": 5,
"disabled": True,
"statetype": "keep state",
}


async def test_kill_states_for_rule_resolves_alias_to_prefix(client):
rule = {"source": "kids", "destination": "any"}
aliases = [{"name": "kids", "type": "network", "address": ["10.0.10.0/24"]}]
with aioresponses() as m:
m.get(f"{API}/firewall/aliases", payload=_envelope(aliases))
m.delete(
re.compile(rf"{re.escape(API)}/firewall/states.*"),
payload=_envelope([]),
repeat=True,
)
assert req.kwargs["json"] == {"id": 5, "disabled": True}
await client.kill_states_for_rule(rule)
deletes = [
str(url)
for (method, url), reqs in m.requests.items()
if method == "DELETE"
for _ in reqs
]
assert any("source__startswith=10.0.10." in u for u in deletes)
assert any("destination__startswith=10.0.10." in u for u in deletes)


async def test_kill_states_for_rule_skips_unresolvable_endpoints(client):
# ``any`` / ``(self)`` / a /25 network have no usable prefix -> no request.
rule = {"source": "any", "destination": "(self)"}
with aioresponses() as m:
await client.kill_states_for_rule(rule)
assert m.requests == {}


async def test_build_telemetry_shape():
Expand Down
Loading