# fastapi_nft_stateless_comment_enums_resilient_full.py """ Stateless FastAPI nftables router with enums and resilient pyroute2 binding. - Stateless: no in-process or on-disk rule store. - Rules may include an optional 'comment' field that will be written into nft's comment. - Uses pyroute2 binding if available and can be instantiated synchronously, otherwise falls back to calling the `nft` CLI via subprocess. - Exposes two endpoints under /nft: GET /nft/rules -> list rules reconstructed from kernel state (returns comment if present) PUT /nft/rules -> replace entire ordered ruleset (clients supply optional comment per rule) Notes: - The service must run with privileges to modify nftables (root / CAP_NET_ADMIN) when applying rules. - If the subprocess fallback is used, ensure the `nft` binary is present. """ from fastapi import APIRouter, HTTPException, Request from pydantic import BaseModel, Field from typing import Optional, List, Dict, Any, Union import logging import json import uuid from enum import Enum import asyncio import subprocess import re from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number # ---------------------- main module (resilient wrapper + API) ---------------------- # Try to import NFTables binding (various pyroute2 layouts) try: from pyroute2.nftables.main import NFTables as NFTablesBinding # type: ignore except Exception: try: from pyroute2.nftables import NFTables as NFTablesBinding # type: ignore except Exception: NFTablesBinding = None # will fall back to subprocess wrapper # Router + logger router = APIRouter() logger = logging.getLogger("nftables") logger.debug("nftables router module loaded") # Defaults DEFAULT_TABLE = "mitm_tbl" DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" # ---------------------- Enums used locally ---------------------- class ActionType(str, Enum): DROP = "drop" ACCEPT = "accept" QUEUE = "queue" REDIRECT = "redirect" class Family(str, Enum): BRIDGE = "bridge" INET = "inet" IP = "ip" IP6 = "ip6" ARP = "arp" # ---------------------- Pydantic models ---------------------- class MatchModel(BaseModel): iif: Optional[str] = None oif: Optional[str] = None meta_length: Optional[Any] = None # int or range string ip_proto: Optional[Union[int, str]] = None # accept number or name (string) tcp_dport: Optional[int] = None udp_dport: Optional[int] = None class ActionModel(BaseModel): type: ActionType queue_num: Optional[int] = None redirect_port: Optional[int] = None class RuleModel(BaseModel): id: Optional[str] = Field(None, description="optional client id; not persisted") family: Optional[Family] = Field(Family.BRIDGE) table: Optional[str] = Field(DEFAULT_TABLE) chain: Optional[str] = Field(DEFAULT_CHAIN) match: MatchModel action: ActionModel comment: Optional[str] = Field(None, description="optional human-readable comment stored in nft comment") class ReplaceResult(BaseModel): applied: bool rules_count: int # ---------------------- resilient NFT wrapper selection ---------------------- class NFTSubprocessWrapper: """Fallback wrapper calling the `nft` CLI via subprocess.""" def __init__(self): try: subprocess.run(["nft", "--version"], check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True) logger.debug("nft CLI available for subprocess wrapper") except Exception as e: logger.warning("nft CLI not available or couldn't be invoked: %s", e) def run(self, cmd: str) -> Dict[str, Any]: args = ["nft"] + cmd.split() try: proc = subprocess.run(args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True) out = proc.stdout or "" try: return json.loads(out) except Exception: return {"out": out} except subprocess.CalledProcessError as e: logger.error("nft CLI failed cmd=%s stderr=%s", cmd, e.stderr) raise RuntimeError(e.stderr or str(e)) def add_table(self, family: str, table: str): return self.run(f"add table {family} {table}") def add_chain(self, family: str, table: str, chain: str, type_: str = "filter", hook: str = "forward", priority: int = 0, policy: str = "accept"): return self.run(f'add chain {family} {table} {chain} {{ type {type_} hook {hook} priority {priority}; policy {policy}; }}') def list_table(self, family: str, table: str): return self.run(f"list table {family} {table}") def list_chain(self, family: str, table: str, chain: str): return self.run(f"list chain {family} {table} {chain} -a") def add_rule(self, family: str, table: str, chain: str, rule_fragment: str): return self.run(f"add rule {family} {table} {chain} {rule_fragment}") def delete_rule_by_handle(self, family: str, table: str, chain: str, handle: str): return self.run(f"delete rule {family} {table} {chain} handle {handle}") class NFTBindingWrapper: """Wrapper using pyroute2 NFTables binding, instantiated synchronously if possible.""" def __init__(self, binding_cls): self._binding = None self._constructed = False tried_kwargs = [ {"async": False}, {"nl_async": False}, {"use_async": False}, {"asyncio": False}, ] last_exc = None for kw in tried_kwargs: try: self._binding = binding_cls(**kw) self._constructed = True logger.debug("NFTables binding instantiated with kwargs %s", kw) break except TypeError as e: last_exc = e except Exception as e: last_exc = e if not self._constructed: try: self._binding = binding_cls() self._constructed = True logger.debug("NFTables binding instantiated with no kwargs") except Exception as e: last_exc = e if not self._constructed: raise RuntimeError(f"failed to instantiate NFTables binding: {last_exc}") setup_coro = getattr(self._binding, "setup_endpoint", None) if setup_coro and asyncio.iscoroutinefunction(setup_coro): if asyncio.get_event_loop().is_running(): raise RuntimeError("pyroute2 NFTables requires async setup but event loop is already running") try: asyncio.get_event_loop().run_until_complete(setup_coro()) except Exception as e: logger.error("awaiting binding.setup_endpoint() failed: %s", e) raise def run(self, cmd: str) -> Dict[str, Any]: try: rc, out, err = self._binding.cmd(cmd) if isinstance(out, bytes): out = out.decode(errors="ignore") if isinstance(err, bytes): err = err.decode(errors="ignore") if rc != 0: raise RuntimeError(err or f"nft cmd failed rc={rc}") if out: try: return json.loads(out) except Exception: return {"out": out} return {} except AttributeError: out = self._binding.json_cmd(cmd) return out or {} except Exception as e: logger.exception("nft binding cmd failed: %s", e) raise def make_nft_wrapper(): if NFTablesBinding is not None: try: if asyncio.get_event_loop().is_running(): logger.info("asyncio loop is running; skipping binding and using subprocess wrapper") raise RuntimeError("event loop running") w = NFTBindingWrapper(NFTablesBinding) logger.info("using pyroute2 NFTables binding") return w except Exception as e: logger.warning("pyroute2 binding unavailable/synchronous construction failed: %s; falling back to nft CLI", e) logger.info("using nft CLI subprocess wrapper") return NFTSubprocessWrapper() NFTC = make_nft_wrapper() logger.info("selected nft wrapper: %s", type(NFTC).__name__) # ---------------------- builders / parsers ---------------------- def build_match_frag(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 is not None: # accept numeric or protocol-name (string). For names try to accept upper or lower. if isinstance(match.ip_proto, int): frag += ["ip", "protocol", str(match.ip_proto)] else: # if it's a known IPProtocolEnum name, use lowercase for nft syntax v = str(match.ip_proto) if v.upper() in IPProtocolEnum.__members__: frag += ["ip", "protocol", v.lower()] else: frag += ["ip", "protocol", v] 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_action_frag(action: ActionModel) -> List[str]: if action.type == ActionType.DROP: return ["drop"] if action.type == ActionType.ACCEPT: return ["accept"] if action.type == ActionType.QUEUE: num = action.queue_num if action.queue_num is not None else 0 return ["queue", "num", str(num)] if action.type == ActionType.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_fragment_from_model(rule: RuleModel) -> str: match_frag = build_match_frag(rule.match) action_frag = build_action_frag(rule.action) parts = match_frag + action_frag if rule.comment: parts += ['comment', f'"{rule.comment}"'] return " ".join(parts) def parse_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 def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]: """ Convert pyroute2/nft `rule` JSON entry into the simple dict format returned by the API. Best-effort parsing. Protocol numbers are converted to IPProtocolEnum member names when possible. """ exprs = [] if "rule" in entry: r = entry["rule"] exprs = r.get("expr") or r.get("expressions") or r.get("exprs") or [] chain_name = r.get("chain") or entry.get("chain") or r.get("chain_name") family = entry.get("family") or DEFAULT_FAMILY table = entry.get("table") or DEFAULT_TABLE else: exprs = entry.get("expr") or entry.get("expressions") or entry.get("exprs") or [] chain_name = entry.get("chain") or entry.get("chain_name") family = entry.get("family") or DEFAULT_FAMILY table = entry.get("table") or DEFAULT_TABLE if isinstance(exprs, dict): exprs = [exprs] comment = parse_comment_from_exprs(exprs) 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") or m.get("field") v = m.get("v") or m.get("s") or m.get("value") if isinstance(v, dict): v = v.get("value") or v.get("v") or v.get("s") if key in ("iifname", "iif", "in"): match["iif"] = v elif key in ("oifname", "oif", "out"): match["oif"] = v elif key in ("length", "len"): match["meta_length"] = v elif "cmp" in e or "match" in e: cmp_obj = e.get("cmp") or e.get("match") or {} left = cmp_obj.get("left") right = cmp_obj.get("right") def _extract_immediate(x): if not x or not isinstance(x, dict): return None for k in ("immediate", "value", "data", "s", "v"): if k in x: val = x[k] if isinstance(val, str) and val.startswith("0x"): try: return int(val, 16) except Exception: return val return val return None imm = _extract_immediate(left) or _extract_immediate(right) # If imm is a string name (like 'icmp'), keep it; if int then map to IPProtocolEnum name if possible. if isinstance(imm, str): low = imm.lower() # If it's a known enum member name, return the enum member name (uppercase) if low.upper() in IPProtocolEnum.__members__: match["ip_proto"] = low.upper() else: # numeric-string? if imm.isdigit(): match["ip_proto"] = protocol_from_number(int(imm)) else: match["ip_proto"] = imm if isinstance(imm, int): if 0 <= imm <= 255: match["ip_proto"] = protocol_from_number(imm) else: # might be a port; treat as tcp_dport if sensible if 0 < imm <= 65535: match.setdefault("tcp_dport", imm) elif "verdict" in e or "immediate" in e or "return" in e: v = e.get("verdict") or e.get("return") or e.get("immediate") if isinstance(v, dict): t = v.get("type") or v.get("kind") or v.get("verdict") if t: action["type"] = t 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 else: # ignore other expression types pass if "type" not in action: action["type"] = "accept" # try to coerce action.type to ActionType try: if isinstance(action.get("type"), str): action["type"] = ActionType(action["type"]) except Exception: pass # coerce family to Family enum if possible try: if isinstance(family, str): family = Family(family) except Exception: pass return { "family": family, "table": table, "chain": chain_name or DEFAULT_CHAIN, "match": match, "action": action, "comment": comment, } # ---------------------- helpers for normalized output ---------------------- def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]: if not out: return [] if isinstance(out, dict): if "nftables" in out and isinstance(out["nftables"], list): return out["nftables"] if "out" in out and isinstance(out["out"], str): text = out["out"] try: parsed = json.loads(text) if isinstance(parsed, dict) and "nftables" in parsed: return parsed["nftables"] if isinstance(parsed, list): return parsed return [parsed] except Exception: return [{"text": text}] return [out] if isinstance(out, list): return out if isinstance(out, str): return [{"text": out}] return [] def _parse_textual_chain_output(text: str, family: str, table: str, chain: str) -> List[Dict[str, Any]]: results: List[Dict[str, Any]] = [] lines = text.splitlines() for raw in lines: line = raw.strip() if not line: continue if line.startswith("type ") or line.startswith("policy ") or line.startswith("chain ") or line.startswith("table "): continue if line in ("{", "}"): continue comment = None cm = re.search(r'comment\s+"([^"]+)"', line) if cm: comment = cm.group(1) line_no_comment = re.sub(r'comment\s+"[^"]+"', '', line) else: line_no_comment = line match: Dict[str, Any] = {} action: Dict[str, Any] = {"type": "accept"} m = re.search(r'\bip\s+protocol\s+([A-Za-z0-9_+-]+)\b', line_no_comment, flags=re.IGNORECASE) if m: proto = m.group(1).lower() if proto.isdigit(): match["ip_proto"] = protocol_from_number(int(proto)) else: # if proto corresponds to an enum member, return its name uppercase, otherwise the raw string if proto.upper() in IPProtocolEnum.__members__: match["ip_proto"] = proto.upper() else: match["ip_proto"] = proto m2 = re.search(r'\btcp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE) if m2: match["tcp_dport"] = int(m2.group(1)) m3 = re.search(r'\budp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE) if m3: match["udp_dport"] = int(m3.group(1)) m4 = re.search(r'\biif\s+"?([^"\s]+)"?\b', line_no_comment, flags=re.IGNORECASE) if m4: match["iif"] = m4.group(1) m5 = re.search(r'\boif\s+"?([^"\s]+)"?\b', line_no_comment, flags=re.IGNORECASE) if m5: match["oif"] = m5.group(1) if re.search(r'\bdrop\b', line_no_comment, flags=re.IGNORECASE): action = {"type": "drop"} elif re.search(r'\baccept\b', line_no_comment, flags=re.IGNORECASE): action = {"type": "accept"} elif re.search(r'\bqueue\b', line_no_comment, flags=re.IGNORECASE): q = re.search(r'queue\s+num\s+(\d+)', line_no_comment, flags=re.IGNORECASE) if q: action = {"type": "queue", "queue_num": int(q.group(1))} else: action = {"type": "queue"} elif re.search(r'\bredirect\b', line_no_comment, flags=re.IGNORECASE) or re.search(r'\bto\s+:\d+\b', line_no_comment, flags=re.IGNORECASE): mredir = re.search(r':(\d+)', line_no_comment) if mredir: action = {"type": "redirect", "redirect_port": int(mredir.group(1))} else: action = {"type": "redirect"} results.append({ "family": family, "table": table, "chain": chain, "match": match, "action": action, "comment": comment, }) return results # ---------------------- high-level operations ---------------------- def ensure_table_chain(family: Union[str, Family], table: str, chain: str): fam = family.value if isinstance(family, Family) else family try: NFTC.add_table(fam, table) except Exception as e: logger.debug("add_table may have failed/exists: %s", e) try: NFTC.add_chain(fam, table, chain) except Exception as e: logger.debug("add_chain may have failed/exists: %s", e) def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEFAULT_CHAIN) -> List[Dict[str, Any]]: fam = family.value if isinstance(family, Family) else family try: raw = NFTC.list_chain(fam, table, chain) logger.debug("raw list_chain output: %s", str(raw)[:2000]) except Exception as e: logger.debug("list_chain failed (%s); falling back to list ruleset", e) try: raw = NFTC.run("list ruleset") logger.debug("raw list ruleset output: %s", str(raw)[:2000]) except Exception as e2: logger.error("list ruleset failed: %s", e2) return [] entries = _normalize_nft_output(raw) results: List[Dict[str, Any]] = [] for item in entries: if not item: continue if isinstance(item, dict) and "text" in item and isinstance(item["text"], str): results.extend(_parse_textual_chain_output(item["text"], fam, table, chain)) continue if "rule" in item: rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")} reconstructed = reconstruct_rule_from_rule_entry(rule_entry) results.append(reconstructed) continue if isinstance(item, dict): if "nftables" in item and isinstance(item["nftables"], list): for sub in item["nftables"]: if isinstance(sub, dict) and "rule" in sub: reconstructed = reconstruct_rule_from_rule_entry(sub) results.append(reconstructed) elif item.get("type") == "rule" or "expr" in item or "expressions" in item: reconstructed = reconstruct_rule_from_rule_entry(item) results.append(reconstructed) else: pass return results def add_rules_replace_all(rules: List[RuleModel], family: Union[str, Family], table: str, chain: str): fam = family.value if isinstance(family, Family) else family try: NFTC.run(f"flush chain {fam} {table} {chain}") except Exception as e: logger.debug("flush chain may have returned error: %s", e) for r in rules: fam_r = r.family.value if isinstance(r.family, Family) else r.family frag = nft_rule_fragment_from_model(r) try: NFTC.run(f"add rule {fam_r} {r.table} {r.chain} {frag}") logger.info("added rule frag=%s", frag) except Exception as e: logger.exception("failed to add rule: %s", e) raise RuntimeError(f"failed to add rule: {e}") # ---------------------- API endpoints ---------------------- @router.get("/rules") def get_rules(family: Optional[Union[str, Family]] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): fam = family.value if isinstance(family, Family) else family logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", fam, table, chain) try: ensure_table_chain(fam, table, chain) except Exception as e: logger.error("failed to ensure table/chain: %s", e) raise HTTPException(status_code=500, detail=str(e)) rules = list_rules_from_nft(fam, table, chain) return {"count": len(rules), "rules": rules} @router.put("/rules", response_model=ReplaceResult) def put_rules( rules: List[RuleModel], request: Request, family: Optional[Union[str, Family]] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN, ): fam = family.value if isinstance(family, Family) else family # validate per-rule family/table/chain for r in rules: r_family_val = r.family.value if isinstance(r.family, Family) else r.family if r.family and r_family_val != fam: raise HTTPException(status_code=400, detail=f"rule family mismatch: {r_family_val} != {fam}") 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 try: ensure_table_chain(fam, table, chain) except Exception as e: logger.error("failed to ensure table/chain: %s", e) raise HTTPException(status_code=500, detail=str(e)) # ensure rule ids for convenience for r in rules: if not r.id: r.id = str(uuid.uuid4()) # attempt to replace rules try: add_rules_replace_all(rules, fam, table, chain) except Exception as e: logger.error("failed to apply rules: %s", e) raise HTTPException(status_code=500, detail=str(e)) logger.info("applied nft ruleset successfully; rules_count=%d", len(rules)) return ReplaceResult(applied=True, rules_count=len(rules))