# fastapi_nft_router.py # -*- coding: utf-8 -*- """ FastAPI router that lists and manages nftables rules for family 'bridge' (default table 'mitm_tbl', chain 'forward'). Behavior: - Uses `nft --json list ruleset` to obtain authoritative rule metadata (handles). - Uses `nft list chain ` to extract the exact textual rule lines. Mapping is done by matching `handle N` in the textual output. - Returns for each rule: - nft_rule_text_full: exact line from 'nft list chain ...' including 'handle N' (or None) - nft_rule_text: same line trimmed to remove trailing 'handle N' (or None) - add_command: "add rule
" (or None) - No JSON->text reconstruction is attempted. If text mapping is missing we return None. Security note: - Process must be run with privileges to run nft (root or appropriate capabilities). - Consider adding auth before exposing these endpoints. """ from typing import Any, Dict, List, Optional, Union, Literal import subprocess import shutil import json import logging import re from enum import Enum from fastapi import APIRouter, HTTPException, Body from pydantic import BaseModel, Field, validator # Router & logging router = APIRouter() logger = logging.getLogger("nftables") logger.debug("nftables router module loaded") # Defaults & nft binary DEFAULT_TABLE = "mitm_tbl" DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" NFT_BIN = shutil.which("nft") # ----------------- Enums (for frontend) ----------------- class Family(str, Enum): bridge = DEFAULT_FAMILY class Table(str, Enum): table = DEFAULT_TABLE class Chain(str, Enum): forward = "forward" input = "input" output = "output" class MetaKey(str, Enum): iifname = "iifname" oifname = "oifname" iif = "iif" oif = "oif" prio = "prio" class EtherField(str, Enum): saddr = "saddr" daddr = "daddr" class IPDir(str, Enum): saddr = "saddr" daddr = "daddr" class Op(str, Enum): eq = "==" neq = "!=" lt = "<" gt = ">" contains = "in" class Verdict(str, Enum): accept = "accept" drop = "drop" reject = "reject" continue_ = "continue" class Proto(str, Enum): tcp = "tcp" udp = "udp" icmp = "icmp" class ConntrackState(str, Enum): new = "new" established = "established" related = "related" invalid = "invalid" class RejectType(str, Enum): icmp = "icmp" tcp_reset = "tcp reset" class IcmpType(str, Enum): dest_unreachable = "destination-unreachable" time_exceeded = "time-exceeded" echo_reply = "echo-reply" echo_request = "echo-request" port_unreachable = "port-unreachable" host_unreachable = "host-unreachable" fragmentation_needed = "fragmentation-needed" class LogGroup(int, Enum): g0 = 0 g1 = 1 g2 = 2 g3 = 3 g4 = 4 g5 = 5 g6 = 6 g7 = 7 # ------------ Pydantic expression models (typed for frontend) ---------- class BaseExpr(BaseModel): kind: str class Config: extra = "forbid" class MetaExpr(BaseExpr): kind: Literal["meta"] = Field(default="meta") key: MetaKey op: Op = Op.eq value: str class EtherExpr(BaseExpr): kind: Literal["ether"] = Field(default="ether") field: EtherField op: Op = Op.eq value: str class IPExpr(BaseExpr): kind: Literal["ip"] = Field(default="ip") side: IPDir op: Op = Op.eq value: str class ProtoPortExpr(BaseExpr): kind: Literal["l4"] = Field(default="l4") proto: Proto sport: Optional[str] = None dport: Optional[str] = None class CTEexpr(BaseExpr): kind: Literal["ct"] = Field(default="ct") state: ConntrackState class VerdictExpr(BaseExpr): kind: Literal["verdict"] = Field(default="verdict") verdict: Verdict class RejectExpr(BaseExpr): kind: Literal["reject"] = Field(default="reject") reject_type: RejectType icmp_type: Optional[IcmpType] = None class LogExpr(BaseExpr): kind: Literal["log"] = Field(default="log") prefix: Optional[str] = None group: Optional[LogGroup] = None class RawExpr(BaseExpr): kind: Literal["raw"] = Field(default="raw") snippet: str Expr = Union[ MetaExpr, EtherExpr, IPExpr, ProtoPortExpr, CTEexpr, VerdictExpr, RejectExpr, LogExpr, RawExpr, ] # ---------------- Rule model ---------------- class RuleModel(BaseModel): family: Family = Family.bridge table: Table = Table.table chain: Chain = Chain.forward expr: List[Expr] = Field(default_factory=list) comment: Optional[str] = None position: Optional[int] = None # 1-based handle: Optional[int] = None @validator("family") def only_bridge(cls, v: Family) -> Family: if v != Family.bridge: raise ValueError("This router only manages family 'bridge'") return v # ----------------- Helpers -------------------- def ensure_nft_available() -> None: if not NFT_BIN: logger.error("nft binary not found on server") raise HTTPException(status_code=500, detail="nft binary not found on server") def run_nft_cmd(cmd: str) -> Dict[str, Any]: """ Execute a single nft script line via `nft -f -`. Returns stdout/stderr. """ ensure_nft_available() full_cmd = [NFT_BIN, "-f", "-"] script = cmd.rstrip() + "\n" logger.info("Running nft command: %s", cmd) logger.debug("Exec: %s ; script: %s", full_cmd, script) try: proc = subprocess.run(full_cmd, input=script.encode(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=True) stdout = proc.stdout.decode() stderr = proc.stderr.decode() logger.info("nft success (stdout %d bytes, stderr %d bytes)", len(stdout), len(stderr)) logger.debug("nft stdout: %s", stdout or "") if stderr: logger.debug("nft stderr: %s", stderr) return {"stdout": stdout, "stderr": stderr} except subprocess.CalledProcessError as e: err = e.stderr.decode() if e.stderr else str(e) logger.error("nft failed: %s", err) raise HTTPException(status_code=500, detail=err) def ensure_table_and_chain_exist(family: str, table: str, chain: str) -> None: """ Ensure the named table and chain exist; create them with conservative defaults if missing. """ logger.debug("Ensure table/chain exist family=%s table=%s chain=%s", family, table, chain) ensure_nft_available() try: out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE) parsed = json.loads(out) except subprocess.CalledProcessError as e: logger.error("Failed to list ruleset: %s", e.stderr.decode()) raise HTTPException(status_code=500, detail=f"nft failed: {e.stderr.decode()}") items = parsed.get("nftables") if isinstance(parsed, dict) else parsed if not isinstance(items, list): items = [] table_exists = False chain_exists = False for it in items: if "table" in it: t = it["table"] if isinstance(t, dict) and t.get("name") == table and t.get("family") == family: table_exists = True if "chain" in it: ch = it["chain"] if isinstance(ch, dict) and ch.get("name") == chain and ch.get("table") == table and ch.get("family") == family: chain_exists = True if not table_exists: logger.info("Creating table %s %s", family, table) run_nft_cmd(f"add table {family} {table}") if not chain_exists: logger.info("Creating chain %s in table %s", chain, table) if chain in ("input", "forward", "output"): run_nft_cmd(f"add chain {family} {table} {chain} {{ type filter hook {chain} priority 0; policy accept; }}") else: run_nft_cmd(f"add chain {family} {table} {chain} {{ policy accept; }}") # ----------------- Text mapping (strict) ----------------- HANDLE_RE = re.compile(r"\bhandle\s+(\d+)\b", flags=re.IGNORECASE) def build_handle_text_map(family: str, table: str, chain: str) -> Dict[int, str]: """ Runs: nft list chain
Returns mapping handle -> full textual line containing 'handle N'. If the textual output cannot be retrieved, raises HTTPException. """ ensure_nft_available() cmd = [NFT_BIN, "--handle", "list", "chain", family, table, chain] logger.debug("Listing chain text: %s", " ".join(cmd)) try: out = subprocess.check_output(cmd, stderr=subprocess.PIPE).decode() except subprocess.CalledProcessError as e: logger.error("Failed to list chain text: %s", e.stderr.decode()) raise HTTPException(status_code=500, detail=e.stderr.decode()) mapping: Dict[int, str] = {} for line in out.splitlines(): s = line.strip() if not s: continue m = HANDLE_RE.search(s) if not m: continue try: h = int(m.group(1)) # full textual line as-is mapping[h] = s logger.debug("Found textual rule for handle %d: %s", h, s) except Exception as ex: logger.debug("Failed parsing handle from line: %s (%s)", s, ex) continue return mapping # ----------------- Rules listing (JSON + strict textual lookup) ----------------- def nft_list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]: """ Return rules parsed from nft --json list ruleset, augmented with textual lines extracted from `nft list chain
` via handle matching. For each rule returned: - family, table, chain - handle - position (1-based in chain) - comment (best-effort from JSON exprs) - verdict (best-effort) - exprs (the JSON expr list) - nft_rule_text_full: exact textual line from nft list chain ... INCLUDING 'handle N' (or None) - nft_rule_text: textual line trimmed to remove trailing 'handle N' (or None) - add_command: "add rule
" (or None) """ ensure_nft_available() # 1) JSON dump: authoritative structure try: out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE) parsed = json.loads(out) except subprocess.CalledProcessError as e: logger.error("Failed to get JSON ruleset: %s", e.stderr.decode()) raise HTTPException(status_code=500, detail=e.stderr.decode()) # 2) textual map: strict mapping by handle text_map: Dict[int, str] = {} try: text_map = build_handle_text_map(DEFAULT_FAMILY, table, chain) logger.debug("Text map size: %d", len(text_map)) except HTTPException as e: # bubble up the error: user asked to extract exact textual lines and we couldn't get them logger.error("Failed to obtain textual chain dump: %s", getattr(e, "detail", str(e))) # still continue — per your request we won't attempt reconstructions, but we can return None textual fields. text_map = {} results: List[Dict[str, Any]] = [] counters: Dict[str, int] = {} items = parsed.get("nftables") if isinstance(parsed, dict) else parsed if not isinstance(items, list): items = [] for it in items: if "rule" not in it: continue r = it["rule"] family = r.get("family") table_name = r.get("table") chain_name = r.get("chain") # only return rules for requested table/chain if table_name != table or chain_name != chain: continue key = f"{family}:{table_name}:{chain_name}" counters.setdefault(key, 0) counters[key] += 1 position = counters[key] handle = r.get("handle") exprs = r.get("expr", []) # best-effort comment + verdict extraction from JSON exprs (keeps UI useful) comment: Optional[str] = None verdict: Optional[str] = None for ex in exprs: if not isinstance(ex, dict): continue if "comment" in ex: c = ex.get("comment") if isinstance(c, str): comment = c elif isinstance(c, dict): comment = c.get("text") or c.get("str") if "verdict" in ex: v = ex["verdict"] if isinstance(v, dict): verdict = next(iter(v.keys()), None) else: verdict = str(v) if "drop" in ex and verdict is None: verdict = "drop" if "accept" in ex and verdict is None: verdict = "accept" if "reject" in ex and verdict is None: verdict = "reject" # strict textual lookup: only use exact line if present in text_map nft_rule_text_full: Optional[str] = None nft_rule_text: Optional[str] = None add_command: Optional[str] = None if handle is not None and handle in text_map: nft_rule_text_full = text_map[handle] # remove trailing ' handle N' to get copy/paste clause m = HANDLE_RE.search(nft_rule_text_full) if m: # slice everything before ' handle N' raw_clause = nft_rule_text_full[: m.start()].strip() else: raw_clause = nft_rule_text_full # remove a trailing lone '#' (and surrounding whitespace) if present # e.g. "meta iifname \"eth0\" # " -> "meta iifname \"eth0\"" nft_rule_text = re.sub(r"\s*#\s*$", "", raw_clause).strip() add_command = f"add rule {table_name} {chain_name} {nft_rule_text}".strip() if nft_rule_text else None logger.debug("Attached textual rule for handle %s", handle) else: logger.debug("No textual mapping for handle %s — textual fields will be None", handle) results.append({ "family": family, "table": table_name, "chain": chain_name, "handle": handle, "position": position, "comment": comment, "verdict": verdict, "exprs": exprs, "nft_rule_text_full": nft_rule_text_full, "nft_rule_text": nft_rule_text, "add_command": add_command, }) return {"rules": results} # ------------ Expr -> nft snippet & command builder (preview/add) ------ def expr_to_nft_snippet(e: Expr) -> str: """Build short nft snippet from typed Expr (used for preview/add).""" if isinstance(e, MetaExpr): val = e.value key = e.key.value return f"meta {key} {e.op.value} {val}" if isinstance(e, EtherExpr): return f"ether {e.field.value} {e.op.value} {e.value}" if isinstance(e, IPExpr): return f"ip {e.side.value} {e.op.value} {e.value}" if isinstance(e, ProtoPortExpr): parts = [e.proto.value] if e.sport: parts.append(f"sport {e.sport}") if e.dport: parts.append(f"dport {e.dport}") return " ".join(parts) if isinstance(e, CTEexpr): return f"ct state {e.state.value}" if isinstance(e, VerdictExpr): return e.verdict.value if e.verdict != Verdict.continue_ else "continue" if isinstance(e, RejectExpr): if e.reject_type == RejectType.icmp: return f"reject with icmp type {e.icmp_type.value}" if e.icmp_type else "reject" return e.reject_type.value if isinstance(e, LogExpr): parts = ["log"] if e.prefix: parts.append(f'prefix "{e.prefix}"') if e.group is not None: parts.append(f"group {int(e.group)}") return " ".join(parts) if isinstance(e, RawExpr): return e.snippet raise ValueError("Unsupported expression type") def rule_to_nft_cmd(rule: RuleModel) -> str: expr_snips = [expr_to_nft_snippet(e) for e in rule.expr] body = " ".join(s for s in expr_snips if s) if rule.position is not None: cmd = f"insert rule {rule.table.value} {rule.chain.value} position {rule.position} {body}" else: cmd = f"add rule {rule.table.value} {rule.chain.value} {body}" if rule.comment: cmd += f' comment "{rule.comment}"' return cmd # ---------------- Endpoints ------------------- @router.get("/options") def get_options() -> Dict[str, Any]: """Return enum choices for frontend dropdowns.""" return { "family": [f.value for f in Family], "table": [t.value for t in Table], "chain": [c.value for c in Chain], "meta_keys": [m.value for m in MetaKey], "ether_fields": [e.value for e in EtherField], "ip_dirs": [d.value for d in IPDir], "ops": [o.value for o in Op], "verdicts": [v.value for v in Verdict], "protocols": [p.value for p in Proto], "ct_states": [s.value for s in ConntrackState], "reject_types": [r.value for r in RejectType], "icmp_types": [i.value for i in IcmpType], "log_groups": [int(g.value) for g in LogGroup], } @router.get("/rules") def list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]: """List rules for the given table/chain (ensures table/chain exist first).""" ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain) return nft_list_rules(table=table, chain=chain) @router.post("/rules/preview") def preview_rule(rule: RuleModel = Body(...)) -> Dict[str, str]: """Return the nft command that would be executed for the provided rule (preview only).""" try: cmd = rule_to_nft_cmd(rule) except Exception as e: logger.error("Preview build failed: %s", e) raise HTTPException(status_code=400, detail=str(e)) return {"cmd": cmd} @router.post("/rules") def add_rule(rule: RuleModel = Body(...)) -> Dict[str, Any]: """Insert/append rule (creates table/chain if missing).""" ensure_table_and_chain_exist(rule.family.value, rule.table.value, rule.chain.value) cmd = rule_to_nft_cmd(rule) return run_nft_cmd(cmd) @router.delete("/rules/{handle}") def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]: """Delete rule by nft handle.""" ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain) cmd = f"delete rule {table} {chain} handle {handle}" return run_nft_cmd(cmd) @router.put("/rules/{handle}") def update_rule(handle: int, rule: RuleModel = Body(...), table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]: """Replace a rule by handle: delete by handle then insert replacement (attempt to preserve position).""" ensure_table_and_chain_exist(rule.family.value, rule.table.value, rule.chain.value) rules_info = nft_list_rules(table=table, chain=chain) position: Optional[int] = None for r in rules_info.get("rules", []): if r.get("handle") == handle: position = r.get("position") break delete_rule(handle, table=table, chain=chain) if position is not None: rule.position = position return add_rule(rule)