diff --git a/daemon/server.py b/daemon/server.py index 57688a4..47f8b60 100644 --- a/daemon/server.py +++ b/daemon/server.py @@ -532,13 +532,21 @@ def main() -> None: loop = asyncio.new_event_loop() - def _on_shutdown(_sig: int) -> None: + async def _shutdown() -> None: + """Graceful shutdown: cancel poller, close runner, teardown.""" logger.info("Shutting down daemon...") _stop_polling() - loop.stop() + try: + await asyncio.wait_for(runner.cleanup(), timeout=5) + except asyncio.TimeoutError: + logger.warning("Runner cleanup timed out, abandoning") + if Path(socket_path).exists(): + os.unlink(socket_path) + logger.info("vacuum-walld stopped") + loop.call_soon(loop.stop) for sig in (signal.SIGTERM, signal.SIGINT): - loop.add_signal_handler(sig, _on_shutdown, sig) + loop.add_signal_handler(sig, lambda: loop.create_task(_shutdown())) runner = web.AppRunner(app) loop.run_until_complete(runner.setup()) @@ -562,10 +570,11 @@ def main() -> None: try: loop.run_forever() finally: - loop.run_until_complete(runner.cleanup()) + # _shutdown() handles cleanup when invoked via signal handler; + # this block is only reached if shutdown didn't happen cleanly + # (e.g., unexpected exit), in which case we unlink the socket. if Path(socket_path).exists(): os.unlink(socket_path) - logger.info("vacuum-walld stopped") if __name__ == "__main__": diff --git a/lib/firewall.py b/lib/firewall.py index 17320ba..9774c6f 100644 --- a/lib/firewall.py +++ b/lib/firewall.py @@ -66,7 +66,9 @@ def _parse_interfaces(output: str) -> list[str]: def _parse_zone_output(zone: str, output: str) -> dict[str, Any]: - """Parse ``firewall-cmd --zone=Z --list-all`` output.""" + """Parse ``firewall-cmd --zone=Z --list-all`` or a zone block + from ``--list-all-zones`` output. + """ info: dict[str, Any] = {"name": zone} for line in output.splitlines(): line = line.strip() @@ -76,6 +78,11 @@ def _parse_zone_output(zone: str, output: str) -> dict[str, Any]: key = key.strip() value = value.strip() + # --list-all-zones uses "rich rules" (space) while + # --zone=Z --list-all uses "rich-rules" (hyphen); normalize. + if key == "rich rules": + key = "rich-rules" + if not value: if key in ("masquerade", "ics"): info[key] = False @@ -116,6 +123,45 @@ def _parse_zone_output(zone: str, output: str) -> dict[str, Any]: return info +def _parse_all_zones_output(output: str) -> dict[str, dict[str, Any]]: + """Parse the combined ``firewall-cmd --list-all-zones`` output. + + Returns a dict mapping each zone name to its parsed info dict + (same structure as ``_parse_zone_output``). + """ + zones: dict[str, dict[str, Any]] = {} + current_name: str | None = None + current_lines: list[str] = [] + + for raw_line in output.splitlines(): + if not raw_line.strip(): + continue + # Non-indented line starts a new zone block + if raw_line[0].isspace(): + if current_name is not None: + current_lines.append(raw_line.strip()) + else: + # Finalize previous zone + if current_name is not None and current_lines: + zones[current_name] = _parse_zone_output( + current_name, "\n".join(current_lines) + ) + # Extract zone name (discard trailing parenthetical metadata) + name = raw_line.strip().split()[0] + if "(" in name: + name = name[: name.index("(")] + current_name = name + current_lines = [] + + # Finalize last zone + if current_name is not None and current_lines: + zones[current_name] = _parse_zone_output( + current_name, "\n".join(current_lines) + ) + + return zones + + # --------------------------------------------------------------------------- # Helpers for parsing forward-port lines # --------------------------------------------------------------------------- @@ -347,6 +393,7 @@ __all__ = [ "_normalize_target", "_now_iso", "_parse_active_zones", + "_parse_all_zones_output", "_parse_forward_ports", "_parse_interfaces", "_parse_zone_output", diff --git a/lib/state.py b/lib/state.py index 8ca3516..507777e 100644 --- a/lib/state.py +++ b/lib/state.py @@ -15,7 +15,7 @@ from typing import Any, ClassVar from lib.common import load_json, run, run_proc from lib.firewall import ( _parse_active_zones, - _parse_zone_output, + _parse_all_zones_output, ) from lib.firewall import ( config_pending as _config_pending, @@ -401,7 +401,6 @@ def _collect_firewall() -> dict[str, Any]: Dict containing firewall zones, interfaces, rules, config, and pending changes. """ - zone_names = run(["firewall-cmd", "--get-zones"], sudo=True).split() active_raw = run(["firewall-cmd", "--get-active-zones"], sudo=True) active = _parse_active_zones(active_raw) services = run(["firewall-cmd", "--get-services"], sudo=True).split() or [] @@ -458,14 +457,13 @@ def _collect_firewall() -> dict[str, Any]: ifaces = list(iface_map.values()) + # Collect all zones in a single call (replaces per-zone loop) zones: dict[str, dict[str, Any]] = {} - for zn in zone_names: - try: - zones[zn] = _parse_zone_output( - zn, run(["firewall-cmd", f"--zone={zn}", "--list-all"], sudo=True) - ) - except Exception: - continue + try: + all_zones_raw = run(["firewall-cmd", "--list-all-zones"], sudo=True) + zones = _parse_all_zones_output(all_zones_raw) + except Exception: + pass # Load config fw_config_path = PROJECT_DIR / "config" / "firewall" / "config.json" @@ -584,7 +582,7 @@ def _collect_dnsmasq() -> dict[str, Any]: "interface": parts[4] if len(parts) > 4 else "", } ) - except RuntimeError: + except Exception: pass # Check config file on disk diff --git a/system/systemd/vacuum-walld.service b/system/systemd/vacuum-walld.service index a9674b5..12d3e27 100644 --- a/system/systemd/vacuum-walld.service +++ b/system/systemd/vacuum-walld.service @@ -12,6 +12,7 @@ WorkingDirectory={{ PROJECT_DIR }} ExecStart={{ PROJECT_DIR }}/.venv/bin/python -m daemon Restart=on-failure RestartSec=5 +TimeoutStopSec=15 Environment=PATH=/usr/local/bin:/usr/bin Environment=PYTHONUNBUFFERED=1 Environment=ACME_HOME={{ PROJECT_DIR }}/data/acme diff --git a/tests/test_firewall.py b/tests/test_firewall.py index d44d7a3..e4102ff 100644 --- a/tests/test_firewall.py +++ b/tests/test_firewall.py @@ -542,3 +542,89 @@ class TestLibParseForwardPorts: def test_empty_string(self): assert firewall._parse_forward_ports("") == [] + + +# --------------------------------------------------------------------------- +# lib/firewall.py — parse all zones output (--list-all-zones) +# --------------------------------------------------------------------------- + + +class TestParseAllZonesOutput: + def test_parses_single_zone(self): + result = firewall._parse_all_zones_output( + "public\n" + " target: default\n" + " interfaces: eth0\n" + " services: ssh http\n" + " masquerade: yes\n" + " rich rules: \n" + ) + assert "public" in result + assert result["public"]["name"] == "public" + assert result["public"]["interfaces"] == ["eth0"] + assert result["public"]["services"] == ["ssh", "http"] + assert result["public"]["masquerade"] is True + assert result["public"]["rich-rules"] == [] + + def test_parses_multiple_zones(self): + result = firewall._parse_all_zones_output( + "public (default, active)\n" + " target: default\n" + " interfaces: eth0\n" + " services: ssh\n" + " masquerade: no\n" + " rich rules: \n" + "internal (active)\n" + " target: ACCEPT\n" + " interfaces: eth1\n" + " services: dhcp\n" + " masquerade: no\n" + " rich rules: \n" + "trusted\n" + " target: ACCEPT\n" + " interfaces: \n" + " services: \n" + " masquerade: no\n" + " rich rules: \n" + ) + assert set(result.keys()) == {"public", "internal", "trusted"} + assert result["public"]["interfaces"] == ["eth0"] + assert result["internal"]["target"] == "ACCEPT" + assert result["trusted"]["services"] == [] + + def test_empty_output(self): + assert firewall._parse_all_zones_output("") == {} + assert firewall._parse_all_zones_output("\n \n") == {} + + def test_handles_blank_lines_between_zones(self): + result = firewall._parse_all_zones_output( + "public\n" + " target: default\n" + " interfaces: eth0\n" + " rich rules: \n" + "\n" + "internal\n" + " target: ACCEPT\n" + " interfaces: eth1\n" + " rich rules: \n" + ) + assert "public" in result + assert "internal" in result + assert result["public"]["interfaces"] == ["eth0"] + assert result["internal"]["interfaces"] == ["eth1"] + + def test_all_default_fields_present(self): + result = firewall._parse_all_zones_output( + "dmz\n" + " target: default\n" + " interfaces: \n" + " services: \n" + " rich rules: \n" + ) + zone = result["dmz"] + for field in ( + "interfaces", "sources", "services", "ports", "protocols", + "forward-ports", "masquerade", "ics", "icmp-blocks", "module", + "target", "rich-rules", + ): + assert field in zone, f"Missing field: {field}" diff --git a/tests/test_state.py b/tests/test_state.py index 8d7fd3b..6d99e80 100644 --- a/tests/test_state.py +++ b/tests/test_state.py @@ -35,8 +35,6 @@ class TestCollectAll: from lib.state import _collect_firewall def run_side(args, **kwargs): - if "--get-zones" in args: - return "public\ninternal" if "--get-active-zones" in args: return "public\n eth0" if "--get-services" in args: @@ -45,9 +43,18 @@ class TestCollectAll: if "link" in args: return "1: lo: mtu 65536\n2: eth0: mtu 1500 link/ether aa:bb\n" return "" - if "--list-all" in args: - return "target: default\ninterfaces: eth0\nsources: \nservices: \nports: \nprotocols: \nforward-ports: \nmasquerade: no\nics: no\nrich-rules: \nicmp-blocks: \nmodule: \n" - return "" + if "--list-all-zones" in args: + return ( + "public\n" + " target: default\n" + " interfaces: eth0\n" + " services: \n" + " ports: \n" + " protocols: \n" + " forward-ports: \n" + " masquerade: no\n" + " rich rules: \n" + ) mock_run.side_effect = run_side result = _collect_firewall() @@ -62,10 +69,8 @@ class TestCollectAll: from lib.state import _collect_firewall def run_side(args, **kwargs): - if "--get-zones" in args: - return "public\ninternal" if "--get-active-zones" in args: - return "public\n eth0\ninternal eth0.100" + return "public\n eth0\ninternal\n eth0.100" if "--get-services" in args: return "ssh http" if "ip" in args[0]: @@ -81,14 +86,27 @@ class TestCollectAll: "3: eth0.100@if100 inet 10.0.0.1/24\n" ) return "" - if "--list-all" in args: + if "--list-all-zones" in args: return ( - "target: default\ninterfaces: eth0\nsources: " - "services: \nports: \nprotocols: \nforward-ports: " - "masquerade: no\nics: no\nrich-rules: " - "icmp-blocks: \nmodule: \n" + "public\n" + " target: default\n" + " interfaces: eth0\n" + " services: \n" + " ports: \n" + " protocols: \n" + " forward-ports: \n" + " masquerade: no\n" + " rich rules: \n" + "internal\n" + " target: ACCEPT\n" + " interfaces: eth0.100\n" + " services: \n" + " ports: \n" + " protocols: \n" + " forward-ports: \n" + " masquerade: no\n" + " rich rules: \n" ) - return "" mock_run.side_effect = run_side result = _collect_firewall() diff --git a/webui/static/pages/proxy.js b/webui/static/pages/proxy.js index a702341..11bd5ba 100644 --- a/webui/static/pages/proxy.js +++ b/webui/static/pages/proxy.js @@ -203,9 +203,7 @@ function pathRow(d, certs, domainPaths, state) { const multiPath = (domainPaths || []).length > 1; let actions; - if (isMgmt) { - actions = Badge({ text: 'mgmt', variant: 'warning' }); - } else if (isWs) { + if (isWs) { actions = ActionButton({ url: '/api/proxy/domains/' + enc(d.domain), method: 'PUT',