"""Unit tests for OPNsense ping / ping_sweep — no real device required. The OPNsense ping API is job-based: a job is created (``set``), started (``start``), produces statistics that are read back via ``search_jobs``, and has to be stopped and removed again. These tests drive that whole lifecycle against a fake ``_get`` / ``_post``. """ import pytest from unittest.mock import MagicMock, patch from napalm_opnsense.opnsense import OPNsenseDriver # --------------------------------------------------------------------------- # Fake API # --------------------------------------------------------------------------- class FakePingAPI: """Minimal stand-in for the OPNsense diagnostics ping controller.""" def __init__(self, stats=None, model_path=("ping", "settings")): #: hostname -> row returned by search_jobs self.stats = stats or {} #: Keys from the model root down to the field level. A real OPNsense #: nests them two deep ({"ping": {"settings": {...}}}); the tests used #: to assume one, which is why nothing caught the payload going to the #: wrong node. self.model_path = tuple(model_path) self.created = [] # payloads passed to set self.started = [] # job ids passed to start self.stopped = [] # job ids passed to stop self.removed = [] # job ids passed to remove self.jobs = {} # job id -> hostname self.search_calls = 0 self._next_id = 0 def get(self, path): if path == "/api/diagnostics/ping/get": fields = { "hostname": "", "fam": {"ip": {"value": "IPv4", "selected": 1}}, "source_address": "", "packetsize": "", "description": "", } return self._nest(fields) if path == "/api/diagnostics/ping/search_jobs": self.search_calls += 1 return {"rows": [self._row(jid, host) for jid, host in self.jobs.items()]} raise AssertionError(f"unexpected GET {path}") def post(self, path, data=None): if path == "/api/diagnostics/ping/set": self.created.append(data) self._next_id += 1 job_id = f"job-{self._next_id}" self.jobs[job_id] = self.fields(data or {}).get("hostname", "") return {"result": "saved", "uuid": job_id} for action, sink in (("start", self.started), ("stop", self.stopped)): if path.startswith(f"/api/diagnostics/ping/{action}/"): sink.append(path.rsplit("/", 1)[1]) return {"status": "ok"} if path.startswith("/api/diagnostics/ping/remove/"): job_id = path.rsplit("/", 1)[1] self.removed.append(job_id) self.jobs.pop(job_id, None) return {"status": "ok"} raise AssertionError(f"unexpected POST {path}") def _nest(self, fields): """Wrap *fields* in the model path, innermost first.""" node = fields for key in reversed(self.model_path): node = {key: node} return node def fields(self, payload): """The field level of a posted payload, or {} if it went to the wrong node.""" node = payload for key in self.model_path: if not isinstance(node, dict) or key not in node: return {} node = node[key] return node if isinstance(node, dict) else {} def _row(self, job_id, hostname): row = {"id": job_id, "hostname": hostname, "status": "running"} row.update(self.stats.get(hostname, {"send": 0, "received": 0})) return row def _driver(fake): """An OPNsenseDriver wired to *fake* instead of a real firewall.""" with patch("napalm_opnsense.opnsense.requests.Session"): drv = OPNsenseDriver(hostname="fw", username="k", password="s") drv.session = MagicMock() drv._get = fake.get drv._post = fake.post return drv def alive_row(rtt=1.5, send=1, received=1): return { "send": send, "received": received, "loss": "0.0 %", "min": rtt, "avg": rtt, "max": rtt, "std-dev": 0.0, "last_error": None, } def dead_row(send=1): return { "send": send, "received": 0, "loss": "100.0 %", "min": 0.0, "avg": 0.0, "max": 0.0, "std-dev": 0.0, "last_error": None, } @pytest.fixture def api(): return FakePingAPI() @pytest.fixture def driver(api): """Driver whose HTTP layer is replaced by the fake ping API.""" with patch("napalm_opnsense.opnsense.requests.Session"): drv = OPNsenseDriver( hostname="fw.example.com", username="api_key", password="api_secret", optional_args={"verify": False}, ) drv.session = MagicMock() drv._get = api.get drv._post = api.post with patch("napalm_opnsense.ping_mixin.time.sleep"): yield drv # --------------------------------------------------------------------------- # ping() # --------------------------------------------------------------------------- def test_ping_reachable_host_returns_napalm_success(driver, api): api.stats["10.0.0.1"] = alive_row(rtt=2.5, send=2, received=2) result = driver.ping("10.0.0.1", count=2) assert result["success"]["probes_sent"] == 2 assert result["success"]["packet_loss"] == 0 assert result["success"]["rtt_avg"] == 2.5 assert result["success"]["results"] == [{"ip_address": "10.0.0.1", "rtt": 2.5}] def test_ping_unreachable_host_reports_full_packet_loss(driver, api): api.stats["10.0.0.9"] = dead_row(send=2) result = driver.ping("10.0.0.9", count=2) assert result["success"]["probes_sent"] == 2 assert result["success"]["packet_loss"] == 2 assert result["success"]["results"] == [] def test_ping_creates_job_with_destination_and_packetsize(driver, api): api.stats["10.0.0.1"] = alive_row() driver.ping("10.0.0.1", size=64, source="10.0.0.254") assert api.fields(api.created[0]) == { "hostname": "10.0.0.1", "fam": "ip", "packetsize": "64", "interval": "1", "source_address": "10.0.0.254", "description": "netork ping sweep", } def test_ping_uses_ipv6_family_for_v6_destination(driver, api): api.stats["2001:db8::1"] = alive_row() driver.ping("2001:db8::1") assert api.fields(api.created[0])["fam"] == "ip6" def test_ping_starts_stops_and_removes_the_job(driver, api): api.stats["10.0.0.1"] = alive_row() driver.ping("10.0.0.1") assert api.started == ["job-1"] assert api.stopped == ["job-1"] assert api.removed == ["job-1"] def test_ping_removes_job_even_when_reading_stats_fails(driver, api): def _boom(path): if path == "/api/diagnostics/ping/search_jobs": raise RuntimeError("API down") return api.get(path) driver._get = _boom result = driver.ping("10.0.0.1") assert "error" in result assert api.removed == ["job-1"] def test_ping_returns_error_when_job_creation_fails(driver, api): api.post = lambda path, data=None: {"result": "failed", "validations": {"hostname": "bad"}} driver._post = api.post result = driver.ping("nope") assert "error" in result assert "hostname" in result["error"] def test_ping_settings_go_to_the_node_the_api_asks_for(): """A real OPNsense nests the ping model two deep — {"ping": {"settings": {...}}} — and rejects a job whose hostname arrives anywhere else with "ping.settings.hostname: A value is required". Posting the fields one level too high made every probe fail, which a sweep reported as a silent network.""" fake = FakePingAPI() fake.stats["10.0.0.1"] = alive_row() drv = _driver(fake) with patch("napalm_opnsense.ping_mixin.time.sleep"): reply = drv.ping("10.0.0.1") assert fake.created[0] == { "ping": { "settings": { "hostname": "10.0.0.1", "fam": "ip", "packetsize": "100", "interval": "1", "description": "netork ping sweep", } } } assert reply["success"]["probes_sent"] == 1 def test_a_model_that_is_only_one_level_deep_still_works(): """Older firmware exposes the fields directly under one node. The path is read from the API rather than assumed, so both shapes work.""" fake = FakePingAPI(model_path=("settings",)) fake.stats["10.0.0.1"] = alive_row() drv = _driver(fake) with patch("napalm_opnsense.ping_mixin.time.sleep"): drv.ping("10.0.0.1") assert list(fake.created[0]) == ["settings"] assert fake.created[0]["settings"]["hostname"] == "10.0.0.1" # --------------------------------------------------------------------------- # ping_sweep() # --------------------------------------------------------------------------- def test_ping_sweep_probes_all_destinations_in_parallel_batches(driver, api): api.stats["10.0.0.1"] = alive_row(rtt=1.0) api.stats["10.0.0.3"] = alive_row(rtt=3.0) result = driver.ping_sweep(["10.0.0.1", "10.0.0.2", "10.0.0.3"]) assert result["scanned"] == 3 assert result["alive_count"] == 2 assert result["entries"] == [ {"ip": "10.0.0.1", "alive": True, "rtt_ms": 1.0}, {"ip": "10.0.0.2", "alive": False, "rtt_ms": None}, {"ip": "10.0.0.3", "alive": True, "rtt_ms": 3.0}, ] def test_ping_sweep_starts_one_job_per_destination_and_cleans_up(driver, api): result = driver.ping_sweep(["10.0.0.1", "10.0.0.2"]) assert len(api.created) == 2 assert sorted(api.started) == ["job-1", "job-2"] assert sorted(api.removed) == ["job-1", "job-2"] assert result["scanned"] == 2 def test_ping_sweep_reads_stats_once_per_batch_not_once_per_host(driver, api): driver.PING_SWEEP_BATCH_SIZE = 8 api.stats.update({f"10.0.0.{i}": alive_row() for i in range(1, 5)}) driver.ping_sweep([f"10.0.0.{i}" for i in range(1, 5)]) assert api.search_calls == 1 def test_ping_sweep_splits_into_batches(driver, api): driver.PING_SWEEP_BATCH_SIZE = 2 api.stats.update({f"10.0.0.{i}": alive_row() for i in range(1, 6)}) result = driver.ping_sweep([f"10.0.0.{i}" for i in range(1, 6)]) assert api.search_calls == 3 # 2 + 2 + 1 assert result["scanned"] == 5 def test_ping_sweep_truncates_at_max_targets(driver, api): result = driver.ping_sweep([f"10.0.0.{i}" for i in range(1, 11)], max_targets=3) assert result["scanned"] == 3 assert result["truncated"] is True assert len(api.created) == 3 def test_ping_sweep_stops_between_batches_when_asked(driver, api): driver.PING_SWEEP_BATCH_SIZE = 2 calls = {"n": 0} def _should_stop(): calls["n"] += 1 return calls["n"] > 1 result = driver.ping_sweep( [f"10.0.0.{i}" for i in range(1, 6)], should_stop=_should_stop, ) assert result["scanned"] == 2 def test_ping_sweep_reports_progress_per_batch(driver, api): driver.PING_SWEEP_BATCH_SIZE = 2 seen = [] driver.ping_sweep( [f"10.0.0.{i}" for i in range(1, 6)], on_progress=lambda done, total: seen.append((done, total)), ) assert seen == [(2, 5), (4, 5), (5, 5)] def test_ping_sweep_marks_hosts_as_errored_when_a_batch_fails(driver, api): def _boom(path): if path == "/api/diagnostics/ping/search_jobs": raise RuntimeError("API down") return api.get(path) driver._get = _boom result = driver.ping_sweep(["10.0.0.1", "10.0.0.2"]) assert result["alive_count"] == 0 assert all(entry["error"] == "API down" for entry in result["entries"]) assert sorted(api.removed) == ["job-1", "job-2"] def test_ping_sweep_on_empty_list_touches_no_api(driver, api): result = driver.ping_sweep([]) assert result == {"entries": [], "scanned": 0, "alive_count": 0, "truncated": False} assert api.created == [] def test_opnsense_driver_reports_ping_support(): assert OPNsenseDriver.supports_ping() is True