From 55e19209991ca6bdad0a83724759b8100ae72f97 Mon Sep 17 00:00:00 2001 From: malmert Date: Sat, 10 Jan 2026 21:26:16 +0100 Subject: [PATCH] improve models and use better nft lib --- backend/requirements.txt | Bin 754 -> 696 bytes backend/src/api/nftables_api.py | 700 ++++++++++++++++---------------- 2 files changed, 341 insertions(+), 359 deletions(-) diff --git a/backend/requirements.txt b/backend/requirements.txt index 0102c523bb3276e7ddab7114131b0590475b7de9..4aee00282ef170e292ce95d76889cf97ca4ec88f 100644 GIT binary patch delta 7 Ocmeywx`TDY4kiE%@dF reconstruct rule objects from nft kernel state (best-effort) -- PUT /nft/rules -> replace entire ordered ruleset (applies via binding, line-by-line) +Stateless FastAPI nftables router using pyroute2.nftables. +- Stateless: no in-process or on-disk rule store. +- Rules may include an optional 'comment' field that will be written into nft's comment. +- Enums introduced for Action.type, ip_proto (common names) and family. +- Endpoints: + - GET /nft/rules -> reconstruct rules from kernel (returns comment if present) + - PUT /nft/rules -> replace entire ordered rule set (clients supply optional comment per rule) """ 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 +from typing import Optional, List, Dict, Any, Union import logging +import json +import uuid +from enum import Enum -# try import libnftables binding -try: - from nftables import Nftables -except Exception: - Nftables = None # clearer error raised when attempting to use binding +# pyroute2 nftables +from pyroute2 import nftables router = APIRouter(prefix="/nft", tags=["nftables"]) logger = logging.getLogger("nftables") -logger.debug("nftables stateless router module loaded") +logger.debug("nftables stateless-comment-enums router loaded") # Defaults and version token DEFAULT_TABLE = "mitm_tbl" @@ -30,30 +30,54 @@ DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" _current_version: Optional[str] = None + +# ---------------------- Enums ---------------------- +class ActionType(str, Enum): + DROP = "drop" + ACCEPT = "accept" + QUEUE = "queue" + REDIRECT = "redirect" + + +class Protocol(str, Enum): + ICMP = "icmp" + TCP = "tcp" + UDP = "udp" + + +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 like "100-200" - ip_proto: Optional[Any] = None # numeric or name + meta_length: Optional[Any] = None # int or range string + # ip_proto may be Protocol enum, integer or arbitrary string (name) + ip_proto: Optional[Union[int, Protocol, str]] = None tcp_dport: Optional[int] = None udp_dport: Optional[int] = None class ActionModel(BaseModel): - type: Optional[str] = None # drop | accept | queue | redirect + type: ActionType # ActionType enum 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) + 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): @@ -62,216 +86,210 @@ class ReplaceResult(BaseModel): 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." - ) +# ---------------------- nft wrapper ---------------------- +class NFT: + def __init__(self): + self.nft = nftables.NFTables() - -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} {{ }}"] + def run(self, cmd: str) -> Dict[str, Any]: + """ + Run a single nft command string via pyroute2.NFTables.cmd. + Returns parsed JSON if possible, otherwise a dict with 'out' textual output. + Raises RuntimeError on failure. + """ 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) ---------------------- -# Replace your previous _reconstruct_rule_from_exprs and list_rules_from_nft with these. - -def _extract_immediate_value(node): - """Return an integer/string value from various immediate/data shapes, or None.""" - if not node or not isinstance(node, dict): - return None - # common keys used by libnftables JSON - for k in ("immediate", "value", "data", "s", "v"): - if k in node: - val = node[k] - # hex strings like "0x00000050" - if isinstance(val, str) and val.startswith("0x"): + rc, out, err = self.nft.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 int(val, 16) + return json.loads(out) except Exception: - return val - return val - # sometimes immediate sits under {"immediate": {"value": ...}} - if "right" in node or "left" in node: - # caller handles left/right structures - return None + return {"out": out} + return {} + except AttributeError: + # Try json_cmd fallback if older pyroute2 version + out = self.nft.json_cmd(cmd) + return out or {} + except Exception as e: + logger.exception("nft wrapper error for cmd=%s: %s", cmd, e) + raise + + 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}") + + +NFTC = NFT() + + +# ---------------------- 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: + # ip_proto can be Protocol enum, int or str + if isinstance(match.ip_proto, Protocol): + frag += ["ip", "protocol", match.ip_proto.value] + elif isinstance(match.ip_proto, int): + frag += ["ip", "protocol", str(match.ip_proto)] + else: + frag += ["ip", "protocol", str(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_action_frag(action: ActionModel) -> List[str]: + # ActionModel.type is an ActionType enum + 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: + """ + Build fragment after 'add rule '. + Includes comment if rule.comment provided. + """ + match_frag = build_match_frag(rule.match) + action_frag = build_action_frag(rule.action) + parts = match_frag + action_frag + if rule.comment: + # include comment as-is (user-provided). Wrap in quotes. + parts += ['comment', f'"{rule.comment}"'] + return " ".join(parts) + + +def parse_comment_from_exprs(exprs: List[Dict[str, Any]]) -> Optional[str]: + """Scan expressions for a comment expression and return its string if present.""" + 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_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]: +def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]: """ - Improved best-effort mapping from nft expression JSON to our RuleModel fields. - This function attempts to match many of the shapes produced by libnftables. + Given a rule entry returned by NFTC.list_table, return a dict that resembles RuleModel + (family/table/chain + match + action + comment if present). This is best-effort parsing. """ + # normalize expressions + 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) + + # best-effort reconstruction of match & action match: Dict[str, Any] = {} action: Dict[str, Any] = {} - # Track last payload hint (if present) so that subsequent cmp can be interpreted - last_payload_hint: Optional[Dict[str, Any]] = None - - for idx, e in enumerate(exprs): - # handle meta (iif/oif/length) + 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") - # sometimes v is dict if isinstance(v, dict): - v = _extract_immediate_value(v) or v.get("s") or v.get("v") - if key in ("iifname", "iif", "in", "iifname?"): - if v: - match["iif"] = v + 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"): - if v: - match["oif"] = v + match["oif"] = v elif key in ("length", "len"): - if v is not None: - match["meta_length"] = v - else: - logger.debug("meta with unknown key: %s value=%s", key, v) - - # store payload hints for later comparisons - elif "payload" in e: - last_payload_hint = e["payload"] - # make it easier for cmp handling: include index - last_payload_hint["_idx"] = idx - - # cmp (compare) expressions: left/right might be payload / immediate structures + 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") - # extract immediate numeric if present - imm_left = _extract_immediate_value(left) - imm_right = _extract_immediate_value(right) - imm = imm_left if imm_left is not None else imm_right - - # If either immediate is a small int treat as ip_proto - if isinstance(imm, int) and 0 < imm < 256: - # prefer to store as number (frontend may show numeric) - match["ip_proto"] = imm - - # If immediate looks like port (1-65535), try to detect target (tcp/udp) - if isinstance(imm, int) and 0 < imm <= 65535: - assigned = False - # heuristics: if left/right contains a payload hint referring to TCP/UDP or 'dport' strings: - for side in (left, right): - if isinstance(side, dict): - # payload form used by libnftables can include 'protocol' or 'field' - pl = side.get("payload") or side.get("left", {}).get("payload") - if isinstance(pl, dict): - protocol_hint = pl.get("protocol") or pl.get("proto") or pl.get("family") - field_hint = pl.get("field") or pl.get("meta") - sh = json.dumps(pl).lower() - if "tcp" in sh or "sport" in sh or "dport" in sh or "th" in sh: - match["tcp_dport"] = imm - assigned = True - break - if "udp" in sh or "udph" in sh or "udp." in sh: - match["udp_dport"] = imm - assigned = True - break - if not assigned: - # fallback: if ip_proto already indicates tcp(6) or udp(17), assign accordingly - proto = match.get("ip_proto") - if proto in (6, "tcp"): - match["tcp_dport"] = imm - elif proto in (17, "udp"): - match["udp_dport"] = imm - else: - # if we can't know, default to tcp_dport as most common case - match.setdefault("tcp_dport", imm) - - # verdict / immediate verdict expressions -> action + 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 isinstance(imm, str): + low = imm.lower() + if low == "icmp": + match["ip_proto"] = 1 + elif low == "tcp": + match["ip_proto"] = 6 + elif low == "udp": + match["ip_proto"] = 17 + else: + try: + match["ip_proto"] = int(imm) + except Exception: + match["ip_proto"] = imm + if isinstance(imm, int): + if 0 < imm < 256: + match["ip_proto"] = imm + elif 0 < imm <= 65535: + match.setdefault("tcp_dport", imm) elif "verdict" in e or "immediate" in e or "return" in e: - # verdict may be string or dict v = e.get("verdict") or e.get("return") or e.get("immediate") - # normalize dict forms if isinstance(v, dict): t = v.get("type") or v.get("kind") or v.get("verdict") if t: action["type"] = t - # redirect/queue shaped differently across versions if "to" in v: action["type"] = "redirect" action["redirect_port"] = v.get("to") @@ -279,103 +297,128 @@ def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]: action["type"] = "queue" action["queue_num"] = v.get("queue") elif isinstance(v, str): - # strings like "accept" or "drop" action["type"] = v - else: - # sometimes verdict is expressed as nested dict under 'verdict': {'kind':'accept'} - if isinstance(e.get("verdict"), dict): - vv = e["verdict"] - action["type"] = vv.get("kind") or vv.get("type") - - # old-style 'match' entries with left/right payload/immediate - elif "match" in e: - m = e["match"] - left = m.get("left") - right = m.get("right") - imm_left = _extract_immediate_value(left) - imm_right = _extract_immediate_value(right) - if isinstance(imm_left, int) and 0 < imm_left < 256: - match["ip_proto"] = imm_left - if isinstance(imm_right, int) and 0 < imm_right < 256: - match["ip_proto"] = imm_right - else: - # unknown expression type: log (DEBUG) for later tuning - logger.debug("unhandled nft expr type: %s", list(e.keys())) + # ignore other expression types + pass - # default action to accept if kernel default or none found if "type" not in action: action["type"] = "accept" - return {"match": match, "action": action} - - -def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: - """ - Reconstruct rules using libnftables JSON output with more robust parsing. - """ - # ensure binding present and request JSON output explicitly - _ensure_binding_available() - nft = Nftables() - nft.set_json_output(True) - rc, out, err = nft.cmd(f"list table {family} {table}") - if rc != 0: - logger.debug("nft list table returned rc=%s err=%s", rc, err) - return [] - + # convert action.type to ActionType enum if possible try: - data = json.loads(out) + action_type_val = action.get("type") + if isinstance(action_type_val, str): + action["type"] = ActionType(action_type_val) + except Exception: + # leave as-is if conversion fails + pass + + # convert 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, + } + + +# ---------------------- 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.error("failed to parse nft JSON output: %s", 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) -> List[Dict[str, Any]]: + fam = family.value if isinstance(family, Family) else family + try: + out = NFTC.list_table(fam, table) + except Exception as e: + logger.debug("list_table failed: %s", e) return [] + entries = [] + if isinstance(out, dict) and "nftables" in out: + entries = out["nftables"] + elif isinstance(out, dict) and "out" in out and isinstance(out["out"], str): + try: + parsed = json.loads(out["out"]) + if isinstance(parsed, dict) and "nftables" in parsed: + entries = parsed["nftables"] + elif isinstance(parsed, list): + entries = parsed + except Exception: + logger.debug("could not parse textual nft output") + entries = [] + elif isinstance(out, list): + entries = out + elif isinstance(out, dict) and out: + entries = [out] + else: + entries = [] + results: List[Dict[str, Any]] = [] - for item in data.get("nftables", []): - if "rule" not in item: + for item in entries: + if not item: continue - rule_obj = item["rule"] - # different versions might use 'expr', 'expressions' or 'expr' - exprs = rule_obj.get("expr") or rule_obj.get("exprs") or rule_obj.get("expressions") or [] - # if exprs is not a list but a dict (some shaped outputs) normalize - if isinstance(exprs, dict): - # sometimes expressions are nested inside a single dict; try to find array keys - for k in ("expr", "expressions", "exprs"): - v = exprs.get(k) - if isinstance(v, list): - exprs = v - break - else: - # give up and wrap - exprs = [exprs] + if "rule" in item: + rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")} + else: + rule_entry = item + reconstructed = reconstruct_rule_from_rule_entry(rule_entry) + results.append(reconstructed) - recon = _reconstruct_rule_from_exprs(exprs) - chain_name = rule_obj.get("chain") or rule_obj.get("chain_name") or rule_obj.get("table") or DEFAULT_CHAIN - # make output match your RuleModel-ish structure (no id) - 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 +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 + # flush chain + try: + NFTC.run(f"flush chain {fam} {table} {chain}") + except Exception as e: + logger.debug("flush chain may have returned error: %s", e) + + # add rules + for r in rules: + # ensure family value converted to string + fam_r = r.family.value if isinstance(r.family, Family) else r.family + frag = nft_rule_fragment_from_model(r) + try: + NFTC.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[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) +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(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}") + 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(family, table) + rules = list_rules_from_nft(fam, table) return {"count": len(rules), "rules": rules, "version": _current_version} @@ -384,107 +427,46 @@ def put_rules( rules: List[RuleModel], request: Request, if_match: Optional[str] = Header(None, alias="If-Match"), - family: Optional[str] = DEFAULT_FAMILY, + family: Optional[Union[str, Family]] = 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 + fam = family.value if isinstance(family, Family) else family + + # optimistic concurrency 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) + # validate per-rule family/table/chain for r in rules: - if r.family and r.family != family: - raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}") + 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 first + # ensure table/chain exist try: - ensure_table_chain(family, table, chain) - except RuntimeError as e: + 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)) - # build script lines (flush + ordered add rules) - script_lines: List[str] = [] - script_lines.append(f"flush chain {family} {table} {chain}") - + # ensure rule ids for client convenience 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 + # attempt to replace rules 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) + 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)) - 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) + _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))