# fastapi_nft_replace.py """ 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) """ from fastapi import APIRouter, HTTPException from pydantic import BaseModel, Field from typing import Optional, List, Dict, Any, Union, Tuple import subprocess import uuid import json import base64 import re import tempfile import os import logging # Router and logger --------------------------------------------------------- router = APIRouter() logger = logging.getLogger("nftables") logger.debug("nftables router module loaded") # Defaults and in-memory version token -------------------------------------- DEFAULT_TABLE = "mitm_tbl" DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" # in-memory version token updated on successful PUT _current_version: Optional[str] = None # ---------------------- Pydantic models ---------------------- class MatchModel(BaseModel): iif: Optional[str] oif: Optional[str] meta_length: Optional[Union[int, str]] # allow "100-200" ip_proto: Optional[str] tcp_dport: Optional[int] udp_dport: Optional[int] class ActionModel(BaseModel): type: str # drop | accept | queue | redirect queue_num: Optional[int] redirect_port: Optional[int] class RuleModel(BaseModel): id: Optional[str] = Field(None, description="optional rule id; generated if missing") 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 # ---------------------- Utilities ---------------------- def run_nft(args: List[str]) -> Tuple[str, str]: """ Run nft with given args. Returns (stdout, stderr). Raises RuntimeError on non-zero exit with stderr included. """ logger.debug("running nft: %s", " ".join(["nft"] + args)) try: proc = subprocess.run(["nft"] + args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True) logger.debug("nft stdout: %s", proc.stdout.strip()) return proc.stdout, proc.stderr except subprocess.CalledProcessError as e: logger.error("nft failed: %s -- %s", " ".join(e.cmd), e.stderr.strip()) raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}") 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: - 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 """ 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: run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"]) 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: # store rule JSON as base64 to avoid quoting/escaping issues inside nft comment j = json.dumps(rule, separators=(",", ":")) b = base64.b64encode(j.encode()).decode() rid = rule.get("id") or "" return f"mitm_id:{rid} mitm_json:{b}" def decode_comment_payload(comment: str) -> Optional[Dict[str, Any]]: # expects comment like: mitm_id: mitm_json: try: parts = comment.split() kv = {p.split(":", 1)[0]: p.split(":", 1)[1] for p in parts if ":" in p} b64 = kv.get("mitm_json") if not b64: return None j = base64.b64decode(b64.encode()).decode() return json.loads(j) except Exception: logger.debug("failed to decode comment payload: %s", comment) return None def build_nft_match_fragment(match: MatchModel) -> List[str]: frag: List[str] = [] if match.iif: frag += ["iif", f'"{match.iif}"'] if match.oif: frag += ["oif", f'"{match.oif}"'] if match.meta_length is not None: frag += ["meta", "length", str(match.meta_length)] if match.ip_proto: frag += ["ip", "protocol", match.ip_proto] if match.tcp_dport: frag += ["tcp", "dport", str(match.tcp_dport)] if match.udp_dport: frag += ["udp", "dport", str(match.udp_dport)] return frag def build_nft_action_fragment(action: ActionModel) -> List[str]: if action.type == "drop": return ["drop"] if action.type == "accept": return ["accept"] if action.type == "queue": num = action.queue_num if action.queue_num is not None else 0 return ["queue", "num", str(num)] if action.type == "redirect": if action.redirect_port is None: raise ValueError("redirect action requires redirect_port") return ["redirect", "to", f":{action.redirect_port}"] raise ValueError(f"unsupported action type: {action.type}") def nft_rule_line_from_model(rule: RuleModel) -> str: """ Produce a single-line nft command: add rule comment "" """ match_frag = build_nft_match_fragment(rule.match) action_frag = build_nft_action_fragment(rule.action) comment = encode_rule_comment(rule.dict()) parts = ["add", "rule", rule.family, rule.table, rule.chain] parts += match_frag parts += action_frag return " ".join(parts) + f' comment "{comment}"' def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: try: out, _ = run_nft(["list", "table", family, table]) except RuntimeError: logger.debug("no table %s.%s found when listing rules", family, table) return [] results: List[Dict[str, Any]] = [] for line in out.splitlines(): line = line.strip() if "comment" in line and "mitm_json:" in line: try: first_quote = line.index('"') last_quote = line.rindex('"') comment_str = line[first_quote + 1:last_quote] except ValueError: comment_str = line.split("comment", 1)[1].strip() parsed = decode_comment_payload(comment_str) if parsed is not None: results.append(parsed) logger.debug("listed %d nft rules from %s.%s", len(results), family, table) return results def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]: ensure_table_chain(rule.family, rule.table, rule.chain) if not rule.id: rule.id = str(uuid.uuid4()) cmd_text = nft_rule_line_from_model(rule) try: # run as discrete args to avoid shell quoting issues run_nft(cmd_text.split()) logger.info("added nft rule id=%s family=%s table=%s chain=%s", rule.id, rule.family, rule.table, rule.chain) return {"id": rule.id, "status": "added"} except RuntimeError as e: logger.error("failed to add rule id=%s: %s", rule.id, e) raise HTTPException(status_code=500, detail=str(e)) def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]: try: out, _ = run_nft(["list", "chain", family, table, chain, "-a"]) except RuntimeError: logger.debug("no chain %s in table %s.%s when attempting delete", chain, family, table) return [] deleted: List[str] = [] for line in out.splitlines(): if "comment" in line and "mitm_id:" in line: try: q1 = line.index('"') q2 = line.index('"', q1 + 1) comment_str = line[q1 + 1:q2] except ValueError: comment_str = line.split("comment", 1)[1] parsed = decode_comment_payload(comment_str) if not parsed: continue rid = parsed.get("id") if rid in ids: m = re.search(r"handle\s+(\d+)", line) if not m: parts = line.split() if "handle" in parts: hi = parts.index("handle") if hi + 1 < len(parts): handle = parts[hi + 1] else: continue else: continue else: handle = m.group(1) try: run_nft(["delete", "rule", family, table, chain, "handle", handle]) logger.info("deleted nft rule id=%s handle=%s", rid, handle) deleted.append(rid) except RuntimeError: logger.error("failed to delete nft rule id=%s handle=%s", rid, handle) continue return deleted # ---------------------- API endpoints ---------------------- @router.get("/rules") def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): """ Ensure the table/chain exist (try to create them if missing), then return the rules from nftables. If ensure_table_chain fails we return 500 with a clear message. """ logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", family, table, chain) try: # Ensure required table/chain exist before listing rules ensure_table_chain(family, table, chain) except RuntimeError as e: logger.error("failed to create/ensure nft table/chain %s.%s/%s: %s", family, table, chain, e) # return an HTTP 500 so the frontend knows setup failed raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {e}") # now safe to list rules rules = list_rules_from_nft(family, table) return {"count": len(rules), "rules": rules, "version": _current_version} @router.post("/rules") def post_rules(payload: Union[RuleModel, List[RuleModel]]): rules = payload if isinstance(payload, list) else [payload] results = [] for r in rules: res = add_rule_to_nft(r) results.append(res) return {"results": results} @router.delete("/rules") def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): ids = [payload] if isinstance(payload, str) else payload deleted = delete_rules_by_ids(ids, family, table, chain) results = [{"id": i, "deleted": i in deleted} for i in ids] return {"results": results} @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): # validate per-rule family/table/chain if present 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 (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: if not r.id: r.id = str(uuid.uuid4()) # build nft script lines lines: List[str] = [] # flush chain (clear existing ordered rules) lines.append(f"flush chain {family} {table} {chain}") # add ordered rules for r in rules: try: lines.append(nft_rule_line_from_model(r)) except Exception as e: raise HTTPException(status_code=400, detail=f"invalid rule: {e}") script = "\n".join(lines) + "\n" tmpfile_path: Optional[str] = None try: with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf: tmpfile_path = tf.name tf.write(script) tf.flush() os.fsync(tf.fileno()) logger.info("wrote nft script to %s; applying...", tmpfile_path) # apply the file try: run_nft(["-f", tmpfile_path]) except RuntimeError as e: logger.error("failed applying nft script: %s", e) raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}") # success: bump version token global _current_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)