improve models and use better nft lib
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s

This commit is contained in:
2026-01-10 21:26:16 +01:00
parent 051087fed3
commit 55e1920999
2 changed files with 341 additions and 359 deletions

Binary file not shown.

View File

@@ -1,28 +1,28 @@
# fastapi_nft_stateless_comment_enums.py
"""
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)
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):
def run(self, cmd: str) -> Dict[str, Any]:
"""
Run a libnftables command string.
Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error.
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.
"""
_ensure_binding_available()
nft = Nftables()
try:
rc, out, err = nft.cmd(cmd)
# ensure we work with strings
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")
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
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:
# 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}")
# ---------------------- chain/table ensure / reconstruction ----------------------
def ensure_table_chain(family: str, table: str, chain: str) -> None:
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:
"""
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.
Build fragment after 'add rule <family> <table> <chain>'.
Includes comment if rule.comment provided.
"""
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)
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)
# ---------------------- 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):
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_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
"""
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] = {}
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
# 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 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
# sometimes immediate sits under {"immediate": {"value": ...}}
if "right" in node or "left" in node:
# caller handles left/right structures
return None
return None
def _reconstruct_rule_from_exprs(exprs: List[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.
"""
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)
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
elif key in ("oifname", "oif", "out"):
if v:
match["oif"] = v
elif key in ("length", "len"):
if v is not None:
match["meta_length"] = v
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:
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
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)
try:
match["ip_proto"] = int(imm)
except Exception:
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
if isinstance(imm, int):
if 0 < imm < 256:
match["ip_proto"] = imm
elif 0 < imm <= 65535:
match.setdefault("tcp_dport", imm)
# verdict / immediate verdict expressions -> action
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")
# ignore other expression types
pass
# 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()))
# 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)
except Exception as e:
logger.error("failed to parse nft JSON output: %s", e)
return []
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
results: List[Dict[str, Any]] = []
for item in data.get("nftables", []):
if "rule" not in 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]
# convert family to Family enum if possible
try:
if isinstance(family, str):
family = Family(family)
except Exception:
pass
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 = {
return {
"family": family,
"table": table,
"chain": chain_name,
"match": recon.get("match", {}),
"action": recon.get("action", {}),
"chain": chain_name or DEFAULT_CHAIN,
"match": match,
"action": action,
"comment": comment,
}
results.append(rule_dict)
logger.debug("reconstructed %d rules from %s.%s", len(results), family, table)
# ---------------------- 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) -> 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 entries:
if not item:
continue
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)
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 "<unknown>", 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)