From 6285e53c2c43666b7efe7b1d4db41a3425f6e97e Mon Sep 17 00:00:00 2001 From: malmert Date: Sat, 10 Jan 2026 18:44:56 +0100 Subject: [PATCH] add initialization for nftables --- backend/src/api/nftables_api.py | 196 ++++++++++++++++++++++++++------ 1 file changed, 164 insertions(+), 32 deletions(-) diff --git a/backend/src/api/nftables_api.py b/backend/src/api/nftables_api.py index 4930efc..18631d0 100644 --- a/backend/src/api/nftables_api.py +++ b/backend/src/api/nftables_api.py @@ -2,16 +2,16 @@ """ nftables router for FastAPI to manage nftables bridge rules dynamically. +This version will attempt to create any kernel/runtime prerequisites: +- load 'bridge' and 'br_netfilter' modules (via modprobe) +- enable sysctls net.bridge.bridge-nf-call-iptables and net.bridge.bridge-nf-call-ip6tables +- create the nft table and chain (with nft add ... and nft -f fallback) + Endpoints (mounted under /nft): - GET /rules -> list active rules (read from nftables) - POST /rules -> add one or many rules (append) - DELETE /rules -> delete one or many rules by id - PUT /rules -> replace entire ordered rule set via nft -f (returns new version) - -Notes: -- Must run with privileges to call `nft` (root or via sudo). -- Rules are stored in the nft rule comment as base64-encoded JSON for round-trip parsing. -- PUT /rules writes a temporary nft script and runs `nft -f `; this applies the new ordered rules. """ from fastapi import APIRouter, HTTPException from pydantic import BaseModel, Field @@ -24,9 +24,10 @@ import re import tempfile import os import logging +import sys # Router and logger --------------------------------------------------------- -router = APIRouter() +router = APIRouter(prefix="/nft", tags=["nftables"]) logger = logging.getLogger("nftables") logger.debug("nftables router module loaded") @@ -87,20 +88,161 @@ def run_nft(args: List[str]) -> Tuple[str, str]: raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}") -def ensure_table_chain(family: str, table: str, chain: str) -> None: - # create table if missing (ignore error if exists) +def _safe_run(cmd: List[str]) -> Tuple[int, str, str]: + """Run arbitrary command and capture exitcode, stdout, stderr. Never raise.""" try: - logger.info("ensuring table %s.%s exists", family, table) - run_nft(["add", "table", family, table]) - except RuntimeError: - logger.debug("table %s.%s may already exist", family, table) + p = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=False) + return p.returncode, p.stdout.strip(), p.stderr.strip() + except Exception as ex: + return 255, "", f"exception: {ex}" - # create chain if missing: forward chain with hook forward + +def _try_modprobe(module: str) -> Tuple[bool, str]: + """Try to modprobe the kernel module. Returns (ok, message).""" + if not shutil_which("modprobe"): + return False, "modprobe not found" + code, out, err = _safe_run(["modprobe", module]) + if code == 0: + return True, out or "ok" + return False, err or f"exit {code}" + + +def shutil_which(cmd: str) -> Optional[str]: + """Small local replacement for shutil.which to avoid extra import in some constrained envs.""" + from shutil import which + return which(cmd) + + +def enable_bridge_sysctls() -> None: + """Attempt to enable sysctls needed for bridge->netfilter interaction.""" + changed = [] + for key, want in [ + ("net.bridge.bridge-nf-call-iptables", "1"), + ("net.bridge.bridge-nf-call-ip6tables", "1"), + ]: + try: + proc = subprocess.run(["/bin/sh", "-c", f"sysctl -w {key}={want}"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=False) + if proc.returncode == 0: + logger.info("set sysctl %s=%s", key, want) + changed.append(key) + else: + # fallback: try write to /proc directly (requires root) + try: + path = "/proc/sys/" + key.replace(".", "/") + if os.path.exists(path): + with open(path, "w") as f: + f.write(want) + logger.info("wrote %s to %s", want, path) + changed.append(key) + else: + logger.warning("sysctl %s not present (and sysctl failed): %s", key, proc.stderr.strip()) + except Exception as e: + logger.warning("failed to write sysctl %s: %s", key, e) + except Exception as e: + logger.debug("sysctl attempt failed for %s: %s", key, e) + if changed: + logger.debug("sysctls changed: %s", changed) + + +def try_ensure_kernel_bridge_support() -> None: + """ + Try to load kernel modules and enable sysctls that are commonly required for 'bridge' family nftables use. + This function logs everything but does not raise. It is best-effort. + """ + logger.debug("attempting to ensure kernel bridge/netfilter support (modprobe + sysctl)") + + # try modprobe bridge and br_netfilter + for mod in ("bridge", "br_netfilter"): + if shutil_which("modprobe"): + code, out, err = _safe_run(["modprobe", mod]) + if code == 0: + logger.info("loaded kernel module: %s", mod) + else: + logger.debug("modprobe %s returned code=%s stderr=%s", mod, code, err) + else: + logger.debug("modprobe not available on this system; skipping module load for %s", mod) + + # enable bridge sysctls so bridged IP packets are seen by netfilter + enable_bridge_sysctls() + + +def ensure_table_chain(family: str, table: str, chain: str) -> None: + """ + Ensure the nft table and chain exist. On serious failures this raises RuntimeError. + This function: + - calls try_ensure_kernel_bridge_support() first (best-effort) + - attempts `nft add table` and `nft add chain` + - if those fail, tries to apply a tiny nft script with `nft -f` to create table and chain + """ + # best-effort kernel prep (modprobe + sysctl) + try: + try_ensure_kernel_bridge_support() + except Exception as e: + logger.debug("kernel prep raised an exception (continuing): %s", e) + + logger.info("ensuring table %s.%s exists", family, table) + + # Try simple add table first + try: + run_nft(["add", "table", family, table]) + logger.debug("created table %s.%s via add table", family, table) + except RuntimeError as e_table: + logger.info("nft add table failed for %s.%s; trying nft -f fallback: %s", family, table, e_table) + # fallback create via nft -f script + script = f"table {family} {table} {{ }}\n" + tmp = None + try: + with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_tbl_create_", suffix=".nft") as tf: + tmp = tf.name + tf.write(script) + tf.flush() + os.fsync(tf.fileno()) + try: + run_nft(["-f", tmp]) + logger.debug("created table %s.%s via nft -f", family, table) + except RuntimeError as e2: + logger.error("nft -f fallback to create table failed: %s (original: %s)", e2, e_table) + raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2 + finally: + if tmp and os.path.exists(tmp): + try: + os.remove(tmp) + except Exception: + pass + + # Now ensure chain exists + logger.info("ensuring chain %s in table %s exists", chain, table) try: - logger.info("ensuring chain %s in table %s exists", chain, table) run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"]) - except RuntimeError: - logger.debug("chain %s in table %s may already exist", chain, table) + logger.debug("created chain %s in %s.%s via add chain", chain, family, table) + except RuntimeError as e_chain: + logger.info("nft add chain failed for %s in %s.%s; trying nft -f fallback: %s", chain, family, table, e_chain) + script = ( + f"table {family} {table} {{\n" + f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n" + f"}}\n" + ) + tmp = None + try: + with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_chain_create_", suffix=".nft") as tf: + tmp = tf.name + tf.write(script) + tf.flush() + os.fsync(tf.fileno()) + try: + run_nft(["-f", tmp]) + logger.debug("created chain %s in %s.%s via nft -f", chain, family, table) + except RuntimeError as e2: + logger.error("nft -f fallback to create chain failed: %s (original: %s)", e2, e_chain) + raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2 + finally: + if tmp and os.path.exists(tmp): + try: + os.remove(tmp) + except Exception: + pass + + logger.info("table/chain ensured: %s.%s/%s", family, table, chain) def encode_rule_comment(rule: Dict[str, Any]) -> str: @@ -289,19 +431,6 @@ def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT @router.put("/rules", response_model=ReplaceResult) def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): - """ - Replace entire ordered rule set by generating an nft script and applying via `nft -f`. - - Behavior: - 1. Validate rules and ensure family/table/chain match (if provided). - 2. Ensure the table/chain exist. - 3. Generate nft script that flushes the chain and adds rules in given order. - 4. Write script to a secure temp file and run `nft -f `. - 5. On success update in-memory version token and return it. - - Note: nft runs script sequentially; if nft errors mid-script, partial state may exist. - For stricter atomicity, implement the temp-table swap approach. - """ # validate per-rule family/table/chain if present for r in rules: if r.family and r.family != family: @@ -311,8 +440,12 @@ def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY, if r.chain and r.chain != chain: raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}") - # ensure table/chain exist - ensure_table_chain(family, table, chain) + # ensure table/chain exist (this will attempt to modprobe + sysctl if needed) + try: + ensure_table_chain(family, table, chain) + except RuntimeError as e: + logger.error("failed to ensure table/chain: %s", e) + raise HTTPException(status_code=500, detail=str(e)) # ensure rule IDs for r in rules: @@ -360,4 +493,3 @@ def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY, os.remove(tmpfile_path) except Exception: logger.debug("failed to remove temp nft script %s", tmpfile_path) -