From 78b687087d341f36306fbb1a738f112e18c45b58 Mon Sep 17 00:00:00 2001 From: malmert Date: Sat, 10 Jan 2026 20:06:09 +0100 Subject: [PATCH] improve --- backend/src/api/nftables_api.py | 305 ++++++++++++++------------------ 1 file changed, 131 insertions(+), 174 deletions(-) diff --git a/backend/src/api/nftables_api.py b/backend/src/api/nftables_api.py index 1a8d967..aeb6d08 100644 --- a/backend/src/api/nftables_api.py +++ b/backend/src/api/nftables_api.py @@ -1,46 +1,41 @@ -# fastapi_nft_stateless.py """ -Stateless FastAPI nftables router. -Only uses kernel-stored info (nft) as the source of truth. +Stateless FastAPI nftables router that uses libnftables binding. -Endpoints (mounted under /nft): -- GET /rules -> reconstruct rule objects from nft kernel state (best-effort) -- PUT /rules -> replace entire ordered ruleset (applies via nft -f) +- GET /nft/rules -> reconstruct rule objects from nft kernel state (best-effort) +- PUT /nft/rules -> replace entire ordered ruleset (applies via binding, line-by-line) -No persistence, no comments used for mapping. """ from fastapi import APIRouter, HTTPException, Header, Request from pydantic import BaseModel, Field -from typing import Optional, List, Dict, Any, Union, Tuple +from typing import Optional, List, Dict, Any import uuid import json import tempfile import os import logging -# libnftables binding +# try import libnftables binding try: from nftables import Nftables except Exception: - Nftables = None # will raise when used + Nftables = None # clearer error raised when attempting to use binding router = APIRouter(prefix="/nft", tags=["nftables"]) logger = logging.getLogger("nftables") -logger.debug("nftables stateless router loaded") +logger.debug("nftables stateless router module loaded") -# Defaults +# Defaults and version token DEFAULT_TABLE = "mitm_tbl" DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" - _current_version: Optional[str] = None # ---------------------- Pydantic models ---------------------- class MatchModel(BaseModel): iif: Optional[str] = None oif: Optional[str] = None - meta_length: Optional[Union[int, str]] = None - ip_proto: Optional[Union[int, str]] = None + meta_length: Optional[Any] = None # int or range string like "100-200" + ip_proto: Optional[Any] = None # numeric or name tcp_dport: Optional[int] = None udp_dport: Optional[int] = None @@ -52,8 +47,8 @@ class ActionModel(BaseModel): class RuleModel(BaseModel): - # note: since we don't persist IDs, id remains optional - id: Optional[str] = Field(None, description="optional client-provided id (not stored by server)") + # id is optional client-side convenience only (not persisted) + id: Optional[str] = Field(None, description="optional client id; not stored server-side") family: Optional[str] = Field(DEFAULT_FAMILY) table: Optional[str] = Field(DEFAULT_TABLE) chain: Optional[str] = Field(DEFAULT_CHAIN) @@ -68,198 +63,162 @@ class ReplaceResult(BaseModel): # ---------------------- nft binding helpers ---------------------- -def _ensure_binding(): +def _ensure_binding_available(): if Nftables is None: raise RuntimeError( - "python nftables binding not available. Install python3-nftables (system package) or pip-nftables." + "python nftables binding not available. Install system package `python3-nftables` or a compatible binding." ) -def nft_cmd(cmd: str) -> Tuple[int, str, str]: - """Run a libnftables command and return (rc, stdout, stderr).""" - _ensure_binding() +def nft_cmd(cmd: str): + """ + Run a libnftables command string. + Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error. + """ + _ensure_binding_available() nft = Nftables() try: rc, out, err = nft.cmd(cmd) - # ensure strings + # ensure we work with strings if isinstance(out, bytes): out = out.decode(errors="ignore") if isinstance(err, bytes): err = err.decode(errors="ignore") return rc, out or "", err or "" except Exception as e: - logger.exception("nft binding call failed for command: %s", cmd) - raise RuntimeError(f"nft binding failed: {e}") + logger.exception("unexpected nft binding exception for cmd=%s", cmd) + raise RuntimeError(f"nft binding error: {e}") def nft_run_or_raise(cmd: str) -> str: + """ + Run command via binding and raise RuntimeError if rc != 0. + Returns stdout string on success. + """ rc, out, err = nft_cmd(cmd) logger.debug("nft cmd: %s -> rc=%s out_len=%d err_len=%d", cmd, rc, len(out), len(err)) if rc != 0: - raise RuntimeError(err or f"nft command {cmd} failed with rc={rc}") + # surface err if available, otherwise a generic message + raise RuntimeError(err or f"nft command '{cmd}' failed (rc={rc})") return out -# ---------------------- reconstruction logic ---------------------- +# ---------------------- chain/table ensure / reconstruction ---------------------- def ensure_table_chain(family: str, table: str, chain: str) -> None: """ - Ensure table and chain exist. Use libnftables commands; fallback to nft -f temp file - if direct add fails. Raises RuntimeError on failure. + Ensure the nft table and chain exist. Use libnftables `add` commands first; on failure + attempt to apply a tiny script by invoking individual commands (no "-f" with the binding). + Raises RuntimeError on persistent failure. """ logger.info("ensuring table %s.%s exists", family, table) + + # try add table try: nft_run_or_raise(f"add table {family} {table}") + logger.debug("created table %s.%s via add table", family, table) except RuntimeError as e_table: - logger.info("add table failed: %s; trying -f fallback", e_table) - script = f"table {family} {table} {{ }}\n" - tmp = None + logger.info("add table failed: %s. attempting fallback", e_table) + # fallback: build script and apply the lines individually + script_lines = [f"table {family} {table} {{ }}"] 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()) - nft_run_or_raise(f"-f {tmp}") - finally: - if tmp and os.path.exists(tmp): - try: - os.remove(tmp) - except Exception: - pass + for ln in script_lines: + nft_run_or_raise(ln) + logger.debug("created table %s.%s via fallback lines", family, table) + except RuntimeError as e2: + logger.error("fallback for creating table failed: %s (original: %s)", e2, e_table) + raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2 - logger.info("ensuring chain %s in table %s", chain, table) + # try add chain + logger.info("ensuring chain %s in table %s exists", chain, table) try: nft_run_or_raise( f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}' ) + logger.debug("created chain %s in %s.%s via add chain", chain, family, table) except RuntimeError as e_chain: - logger.info("add chain failed: %s; trying -f fallback", e_chain) - script = ( - f"table {family} {table} {{\n" - f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n" - f"}}\n" - ) - tmp = None + logger.info("add chain failed: %s. attempting fallback", e_chain) + script_lines = [ + f"table {family} {table} {{", + f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}", + f"}}", + ] 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()) - nft_run_or_raise(f"-f {tmp}") - finally: - if tmp and os.path.exists(tmp): - try: - os.remove(tmp) - except Exception: - pass - - -def _extract_comment_from_exprs(exprs: List[Dict[str, Any]]) -> Optional[str]: - for e in exprs: - if "comment" in e: - cm = e["comment"] - if isinstance(cm, dict): - return cm.get("string") or cm.get("s") or cm.get("value") - elif isinstance(cm, str): - return cm - return None + for ln in script_lines: + nft_run_or_raise(ln) + logger.debug("created chain %s in %s.%s via fallback lines", chain, family, table) + except RuntimeError as e2: + logger.error("fallback for creating chain failed: %s (original: %s)", e2, e_chain) + raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2 + + logger.info("table/chain ensured: %s.%s/%s", family, table, chain) +# ---------------------- reconstruct rules from nft JSON (best-effort) ---------------------- def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]: """ - Best-effort mapping from nft expression json to our RuleModel-like dict. - This is heuristic: nft JSON shapes differ across kernel/libnftables versions. - We cover common patterns: meta (iif/oif/length), payload/cmp for ip proto and ports, verdict for action. + Heuristic mapping of nft expression JSON to our RuleModel fields. + Covers common shapes: meta (iif/oif/length), cmp/payload (ip proto and ports), verdict (action). + This is best-effort; complex expressions may not map perfectly. """ match: Dict[str, Any] = {} action: Dict[str, Any] = {} - # iterate expressions; keep simple heuristics for e in exprs: if "meta" in e: m = e["meta"] - # keys differ; check for common names key = m.get("key") or m.get("type") + v = m.get("v") or m.get("s") or m.get("value") if key in ("iifname", "iif"): - v = m.get("v") or m.get("s") or m.get("value") if v: match["iif"] = v elif key in ("oifname", "oif"): - v = m.get("v") or m.get("s") or m.get("value") if v: match["oif"] = v elif key == "length": - v = m.get("v") or m.get("s") or m.get("value") if v is not None: match["meta_length"] = v elif "payload" in e: - # payload indicates reading bytes of header; usually followed by a 'cmp' comparing to immediate - # store payload description to use when we see a cmp - # flatten payload into a marker for later cmp detection - e_payload = e["payload"] - e["_seen_payload"] = e_payload + # payload describes read of header bytes; often next 'cmp' compares it + # we tag the payload on the expression to help cmp heuristics + e["_payload_hint"] = e["payload"] elif "cmp" in e: cmp = e["cmp"] - # cmp can have 'left'/'right' or 'data' fields left = cmp.get("left") right = cmp.get("right") - # helper to pull immediate numeric value def _extract_immediate(node): - if not node: + if not node or not isinstance(node, dict): return None - if isinstance(node, dict): - for k in ("immediate", "value", "data", "s", "v"): - if k in node: - val = node[k] - # hex string -> int - if isinstance(val, str) and val.startswith("0x"): - try: - return int(val, 16) - except Exception: - return val - return val + for k in ("immediate", "value", "data", "s", "v"): + if k in node: + val = node[k] + if isinstance(val, str) and val.startswith("0x"): + try: + return int(val, 16) + except Exception: + return val + return val return None - imm_left = _extract_immediate(left) imm_right = _extract_immediate(right) - - # If either immediate looks like small numeric, assume ip_proto + # ip proto numeric likely in 1..255 for imm in (imm_left, imm_right): if isinstance(imm, int) and 0 < imm < 256: - # set ip_proto numeric match["ip_proto"] = imm break - - # If immediate value looks like a TCP/UDP port (typical 1..65535) and payload context indicates tcp/udp dport, - # it's hard to be 100% sure; we use heuristic: if cmp mentions 'dport' or payload had offset consistent with port, - # or imm is in port range but no ip_proto set, we attempt to set tcp_dport/udp_dport. + # port heuristics (1..65535) imm = imm_left if imm_left is not None else imm_right if isinstance(imm, int) and 0 < imm <= 65535: - # if we previously saw payload indicating tcp or udp, detect from payload description - # naive heuristic: if any seen payload mentions 'tcp' or 'udp' in its dict then assign accordingly - # otherwise assign tcp_dport by default (best-effort) - assigned = False - for ev in (left, right): - if isinstance(ev, dict): - # look for hints - if "payload" in ev: - pd = ev["payload"] - if isinstance(pd, dict) and ("tcp" in str(pd).lower()): - match["tcp_dport"] = imm - assigned = True - break - if not assigned: - # fallback to tcp_dport heuristic - match.setdefault("tcp_dport", imm) + # best-effort assign to tcp_dport (common case) + # more advanced heuristics could inspect payload hints + match.setdefault("tcp_dport", imm) elif "verdict" in e: v = e["verdict"] - # typical shapes: {"verdict":"accept"} or {"verdict":{"type":"drop"}} + # shapes vary: dict or string if isinstance(v, dict): t = v.get("type") or v.get("kind") if t: action["type"] = t - # redirect/queue handling may vary; try to extract + # redirect / queue extra fields vary if "to" in v: action["type"] = "redirect" action["redirect_port"] = v.get("to") @@ -268,25 +227,18 @@ def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]: action["queue_num"] = v.get("queue") elif isinstance(v, str): action["type"] = v - elif "jump" in e: - # jump is effectively a control flow; not modeled - pass - elif "immediate" in e: - # sometimes immediate verdicts - pass - # there are many other expression types; above covers common ones + # other expression types intentionally ignored for stateless reconstruction - # Default action if none found if "type" not in action: - action["type"] = "accept" # kernel often has policy accept if not specified - + # if kernel default, we assume accept + action["type"] = "accept" return {"match": match, "action": action} def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: """ - Reconstruct a list of rules from the nft kernel listing (JSON). - Returns ordered list of dicts each matching RuleModel shape (without id). + Reconstruct an ordered list of rule dicts from `nft list table` JSON via the binding. + Returns list of dicts shaped like RuleModel (without id). """ try: out = nft_run_or_raise(f"list table {family} {table}") @@ -294,35 +246,33 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: logger.debug("no table %s.%s found when listing rules", family, table) return [] - # libnftables returns spammy text sometimes; prefer JSON output - # attempt to parse as JSON if lib produced JSON; otherwise fallback to textual parsing + # attempt to parse JSON output (binding can return textual JSON) try: - # if out is JSON text produced by libnftables, it is already JSON representation of ruleset data = json.loads(out) except Exception: - # fallback: run with JSON mode by using the binding directly to request JSON output - _ensure_binding() + # force JSON mode in binding if previous parsing failed + _ensure_binding_available() nft = Nftables() nft.set_json_output(True) rc, out_json, err = nft.cmd(f"list table {family} {table}") if rc != 0: - logger.debug("nft JSON list failed: %s", err) + logger.debug("nft JSON listing failed: %s", err) return [] data = json.loads(out_json) results: List[Dict[str, Any]] = [] - # nft JSON structure: {"nftables":[ { "table":...}, { "chain":...}, { "rule": {...} }, ... ]} for item in data.get("nftables", []): if "rule" not in item: continue rule_obj = item["rule"] - exprs = rule_obj.get("expr", []) or rule_obj.get("expr", []) # different bindings key names + exprs = rule_obj.get("expr", []) or rule_obj.get("expressions", []) or [] recon = _reconstruct_rule_from_exprs(exprs) - # create a RuleModel-like dict + # the chain name may be in the rule metadata + chain_name = rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN rule_dict = { "family": family, "table": table, - "chain": rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN, + "chain": chain_name, "match": recon.get("match", {}), "action": recon.get("action", {}), } @@ -335,11 +285,10 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: @router.get("/rules") def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", family, table, chain) - # ensure chain exists (create if missing) to provide consistent output try: ensure_table_chain(family, table, chain) except RuntimeError as e: - logger.error("failed to ensure table/chain: %s", e) + logger.error("failed to ensure table/chain %s.%s/%s: %s", family, table, chain, e) raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {e}") rules = list_rules_from_nft(family, table) @@ -356,15 +305,15 @@ def put_rules( chain: Optional[str] = DEFAULT_CHAIN, ): """ - Replace entire ordered rule set. This implementation builds nft commands (without comments) - from the provided rules and applies via 'nft -f' using libnftables fallback. + Replace entire ordered ruleset. Validates provided rules, then builds nft commands + and applies them line-by-line via libnftables binding (no '-f' token passed into binding). """ # optimistic concurrency global _current_version if if_match is not None and _current_version is not None and if_match != _current_version: raise HTTPException(status_code=409, detail="version mismatch; fetch latest rules and retry") - # validate and ensure table/chain exist + # validate per-rule family/table/chain (quick checks) for r in rules: if r.family and r.family != family: raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}") @@ -373,21 +322,24 @@ def put_rules( if r.chain and r.chain != chain: raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}") + # ensure table/chain exist first 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)) - # Build nft script lines (no comments) + # build script lines (flush + ordered add rules) script_lines: List[str] = [] script_lines.append(f"flush chain {family} {table} {chain}") + for r in rules: - # ensure id exists only for client convenience; not stored + # assign client-side id if missing (not stored) if not r.id: r.id = str(uuid.uuid4()) - # build add rule command using the same helpers - match_frag = [] + + # build match fragments + match_frag: List[str] = [] if r.match.iif: match_frag += ["iif", f'"{r.match.iif}"'] if r.match.oif: @@ -401,43 +353,48 @@ def put_rules( if r.match.udp_dport: match_frag += ["udp", "dport", str(r.match.udp_dport)] - action_frag = [] - if r.action.type == "drop": + # build action fragment + action_frag: List[str] = [] + a_type = r.action.type or "accept" + if a_type == "drop": action_frag = ["drop"] - elif r.action.type == "accept" or not r.action.type: + elif a_type == "accept": action_frag = ["accept"] - elif r.action.type == "queue": + elif a_type == "queue": num = r.action.queue_num if r.action.queue_num is not None else 0 action_frag = ["queue", "num", str(num)] - elif r.action.type == "redirect": + elif a_type == "redirect": if r.action.redirect_port is None: raise HTTPException(status_code=400, detail="redirect action requires redirect_port") action_frag = ["redirect", "to", f":{r.action.redirect_port}"] else: - raise HTTPException(status_code=400, detail=f"unsupported action type: {r.action.type}") + raise HTTPException(status_code=400, detail=f"unsupported action type: {a_type}") - parts = ["add", "rule", r.family, r.table, r.chain] - parts += match_frag - parts += action_frag + parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag script_lines.append(" ".join(parts)) - script = "\n".join(script_lines) + "\n" + # apply script lines via binding (line-by-line) tmpfile_path: Optional[str] = None try: + # keep a temp file copy for debugging if desired (optional) with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf: tmpfile_path = tf.name - tf.write(script) + tf.write("\n".join(script_lines) + "\n") tf.flush() os.fsync(tf.fileno()) - logger.info("wrote nft script to %s; applying...", tmpfile_path) + logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path) try: - # use libnftables binding to apply file if possible, otherwise lib will call nft -f underneath - nft_run_or_raise(f"-f {tmpfile_path}") + for ln in script_lines: + ln = ln.strip() + if not ln: + continue + nft_run_or_raise(ln) except RuntimeError as e: - logger.error("failed applying nft script: %s", e) + logger.error("failed applying nft script line '%s': %s", ln if 'ln' in locals() else "", e) raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}") + # success -> bump version _current_version = str(uuid.uuid4()) logger.info("applied nft ruleset successfully; version=%s", _current_version) return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))