""" Stateless FastAPI nftables router that uses libnftables binding. - 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) """ from fastapi import APIRouter, HTTPException, Header, Request from pydantic import BaseModel, Field from typing import Optional, List, Dict, Any import uuid import json import tempfile import os import logging # try import libnftables binding try: from nftables import Nftables except Exception: 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 module loaded") # 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[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 class ActionModel(BaseModel): type: Optional[str] = None # drop | accept | queue | redirect queue_num: Optional[int] = None redirect_port: Optional[int] = None class RuleModel(BaseModel): # 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) match: MatchModel action: ActionModel class ReplaceResult(BaseModel): version: str applied: bool rules_count: int # ---------------------- nft binding helpers ---------------------- def _ensure_binding_available(): if Nftables is None: raise RuntimeError( "python nftables binding not available. Install system package `python3-nftables` or a compatible 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 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("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: # surface err if available, otherwise a generic message raise RuntimeError(err or f"nft command '{cmd}' failed (rc={rc})") return out # ---------------------- chain/table ensure / reconstruction ---------------------- def ensure_table_chain(family: str, table: str, chain: str) -> None: """ 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. attempting fallback", e_table) # fallback: build script and apply the lines individually script_lines = [f"table {family} {table} {{ }}"] try: 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 # 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. attempting fallback", e_chain) script_lines = [ f"table {family} {table} {{", f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}", f"}}", ] try: 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]: """ 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] = {} for e in exprs: if "meta" in e: m = e["meta"] 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"): if v: match["iif"] = v elif key in ("oifname", "oif"): if v: match["oif"] = v elif key == "length": if v is not None: match["meta_length"] = v elif "payload" in e: # 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"] left = cmp.get("left") right = cmp.get("right") def _extract_immediate(node): if not node or not isinstance(node, dict): return None 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) # ip proto numeric likely in 1..255 for imm in (imm_left, imm_right): if isinstance(imm, int) and 0 < imm < 256: match["ip_proto"] = imm break # port heuristics (1..65535) imm = imm_left if imm_left is not None else imm_right if isinstance(imm, int) and 0 < imm <= 65535: # 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"] # shapes vary: dict or string if isinstance(v, dict): t = v.get("type") or v.get("kind") if t: action["type"] = t # redirect / queue extra fields vary if "to" in v: action["type"] = "redirect" action["redirect_port"] = v.get("to") if "queue" in v: action["type"] = "queue" action["queue_num"] = v.get("queue") elif isinstance(v, str): action["type"] = v # other expression types intentionally ignored for stateless reconstruction if "type" not in action: # 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 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}") except RuntimeError: logger.debug("no table %s.%s found when listing rules", family, table) return [] # attempt to parse JSON output (binding can return textual JSON) try: data = json.loads(out) except Exception: # 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 listing failed: %s", err) return [] data = json.loads(out_json) results: List[Dict[str, Any]] = [] 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("expressions", []) or [] recon = _reconstruct_rule_from_exprs(exprs) # 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": chain_name, "match": recon.get("match", {}), "action": recon.get("action", {}), } results.append(rule_dict) logger.debug("reconstructed %d rules from %s.%s", len(results), family, table) return results # ---------------------- API endpoints ---------------------- @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) try: ensure_table_chain(family, table, chain) except RuntimeError as 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) return {"count": len(rules), "rules": rules, "version": _current_version} @router.put("/rules", response_model=ReplaceResult) def put_rules( rules: List[RuleModel], request: Request, if_match: Optional[str] = Header(None, alias="If-Match"), family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN, ): """ 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 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}") if r.table and r.table != table: raise HTTPException(status_code=400, detail=f"rule table mismatch: {r.table} != {table}") 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 script lines (flush + ordered add rules) script_lines: List[str] = [] script_lines.append(f"flush chain {family} {table} {chain}") for r in rules: # assign client-side id if missing (not stored) if not r.id: r.id = str(uuid.uuid4()) # build match fragments match_frag: List[str] = [] if r.match.iif: match_frag += ["iif", f'"{r.match.iif}"'] if r.match.oif: match_frag += ["oif", f'"{r.match.oif}"'] if r.match.meta_length is not None: match_frag += ["meta", "length", str(r.match.meta_length)] if r.match.ip_proto: match_frag += ["ip", "protocol", str(r.match.ip_proto)] if r.match.tcp_dport: match_frag += ["tcp", "dport", str(r.match.tcp_dport)] if r.match.udp_dport: match_frag += ["udp", "dport", str(r.match.udp_dport)] # build action fragment action_frag: List[str] = [] a_type = r.action.type or "accept" if a_type == "drop": action_frag = ["drop"] elif a_type == "accept": action_frag = ["accept"] 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 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: {a_type}") parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag script_lines.append(" ".join(parts)) # 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("\n".join(script_lines) + "\n") tf.flush() os.fsync(tf.fileno()) logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path) try: 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 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)) finally: if tmpfile_path and os.path.exists(tmpfile_path): try: os.remove(tmpfile_path) except Exception: logger.debug("failed to remove temp nft script %s", tmpfile_path)