improve
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s

This commit is contained in:
2026-01-10 20:06:09 +01:00
parent 09d64e6eec
commit 78b687087d

View File

@@ -1,46 +1,41 @@
# fastapi_nft_stateless.py
""" """
Stateless FastAPI nftables router. Stateless FastAPI nftables router that uses libnftables binding.
Only uses kernel-stored info (nft) as the source of truth.
Endpoints (mounted under /nft): - GET /nft/rules -> reconstruct rule objects from nft kernel state (best-effort)
- GET /rules -> reconstruct rule objects from nft kernel state (best-effort) - PUT /nft/rules -> replace entire ordered ruleset (applies via binding, line-by-line)
- PUT /rules -> replace entire ordered ruleset (applies via nft -f)
No persistence, no comments used for mapping.
""" """
from fastapi import APIRouter, HTTPException, Header, Request from fastapi import APIRouter, HTTPException, Header, Request
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any, Union, Tuple from typing import Optional, List, Dict, Any
import uuid import uuid
import json import json
import tempfile import tempfile
import os import os
import logging import logging
# libnftables binding # try import libnftables binding
try: try:
from nftables import Nftables from nftables import Nftables
except Exception: except Exception:
Nftables = None # will raise when used Nftables = None # clearer error raised when attempting to use binding
router = APIRouter(prefix="/nft", tags=["nftables"]) router = APIRouter(prefix="/nft", tags=["nftables"])
logger = logging.getLogger("nftables") logger = logging.getLogger("nftables")
logger.debug("nftables stateless router loaded") logger.debug("nftables stateless router module loaded")
# Defaults # Defaults and version token
DEFAULT_TABLE = "mitm_tbl" DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward" DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge" DEFAULT_FAMILY = "bridge"
_current_version: Optional[str] = None _current_version: Optional[str] = None
# ---------------------- Pydantic models ---------------------- # ---------------------- Pydantic models ----------------------
class MatchModel(BaseModel): class MatchModel(BaseModel):
iif: Optional[str] = None iif: Optional[str] = None
oif: Optional[str] = None oif: Optional[str] = None
meta_length: Optional[Union[int, str]] = None meta_length: Optional[Any] = None # int or range string like "100-200"
ip_proto: Optional[Union[int, str]] = None ip_proto: Optional[Any] = None # numeric or name
tcp_dport: Optional[int] = None tcp_dport: Optional[int] = None
udp_dport: Optional[int] = None udp_dport: Optional[int] = None
@@ -52,8 +47,8 @@ class ActionModel(BaseModel):
class RuleModel(BaseModel): class RuleModel(BaseModel):
# note: since we don't persist IDs, id remains optional # id is optional client-side convenience only (not persisted)
id: Optional[str] = Field(None, description="optional client-provided id (not stored by server)") id: Optional[str] = Field(None, description="optional client id; not stored server-side")
family: Optional[str] = Field(DEFAULT_FAMILY) family: Optional[str] = Field(DEFAULT_FAMILY)
table: Optional[str] = Field(DEFAULT_TABLE) table: Optional[str] = Field(DEFAULT_TABLE)
chain: Optional[str] = Field(DEFAULT_CHAIN) chain: Optional[str] = Field(DEFAULT_CHAIN)
@@ -68,198 +63,162 @@ class ReplaceResult(BaseModel):
# ---------------------- nft binding helpers ---------------------- # ---------------------- nft binding helpers ----------------------
def _ensure_binding(): def _ensure_binding_available():
if Nftables is None: if Nftables is None:
raise RuntimeError( raise RuntimeError(
"python nftables binding not available. Install python3-nftables (system package) or pip-nftables." "python nftables binding not available. Install system package `python3-nftables` or a compatible binding."
) )
def nft_cmd(cmd: str) -> Tuple[int, str, str]: def nft_cmd(cmd: str):
"""Run a libnftables command and return (rc, stdout, stderr).""" """
_ensure_binding() Run a libnftables command string.
Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error.
"""
_ensure_binding_available()
nft = Nftables() nft = Nftables()
try: try:
rc, out, err = nft.cmd(cmd) rc, out, err = nft.cmd(cmd)
# ensure strings # ensure we work with strings
if isinstance(out, bytes): if isinstance(out, bytes):
out = out.decode(errors="ignore") out = out.decode(errors="ignore")
if isinstance(err, bytes): if isinstance(err, bytes):
err = err.decode(errors="ignore") err = err.decode(errors="ignore")
return rc, out or "", err or "" return rc, out or "", err or ""
except Exception as e: except Exception as e:
logger.exception("nft binding call failed for command: %s", cmd) logger.exception("unexpected nft binding exception for cmd=%s", cmd)
raise RuntimeError(f"nft binding failed: {e}") raise RuntimeError(f"nft binding error: {e}")
def nft_run_or_raise(cmd: str) -> str: 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) 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)) logger.debug("nft cmd: %s -> rc=%s out_len=%d err_len=%d", cmd, rc, len(out), len(err))
if rc != 0: if rc != 0:
raise RuntimeError(err or f"nft command {cmd} failed with rc={rc}") # surface err if available, otherwise a generic message
raise RuntimeError(err or f"nft command '{cmd}' failed (rc={rc})")
return out return out
# ---------------------- reconstruction logic ---------------------- # ---------------------- chain/table ensure / reconstruction ----------------------
def ensure_table_chain(family: str, table: str, chain: str) -> None: def ensure_table_chain(family: str, table: str, chain: str) -> None:
""" """
Ensure table and chain exist. Use libnftables commands; fallback to nft -f temp file Ensure the nft table and chain exist. Use libnftables `add` commands first; on failure
if direct add fails. Raises RuntimeError 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) logger.info("ensuring table %s.%s exists", family, table)
# try add table
try: try:
nft_run_or_raise(f"add table {family} {table}") 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: except RuntimeError as e_table:
logger.info("add table failed: %s; trying -f fallback", e_table) logger.info("add table failed: %s. attempting fallback", e_table)
script = f"table {family} {table} {{ }}\n" # fallback: build script and apply the lines individually
tmp = None script_lines = [f"table {family} {table} {{ }}"]
try: try:
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_tbl_create_", suffix=".nft") as tf: for ln in script_lines:
tmp = tf.name nft_run_or_raise(ln)
tf.write(script) logger.debug("created table %s.%s via fallback lines", family, table)
tf.flush() except RuntimeError as e2:
os.fsync(tf.fileno()) logger.error("fallback for creating table failed: %s (original: %s)", e2, e_table)
nft_run_or_raise(f"-f {tmp}") raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2
finally:
if tmp and os.path.exists(tmp):
try:
os.remove(tmp)
except Exception:
pass
logger.info("ensuring chain %s in table %s", chain, table) # try add chain
logger.info("ensuring chain %s in table %s exists", chain, table)
try: try:
nft_run_or_raise( nft_run_or_raise(
f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}' 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: except RuntimeError as e_chain:
logger.info("add chain failed: %s; trying -f fallback", e_chain) logger.info("add chain failed: %s. attempting fallback", e_chain)
script = ( script_lines = [
f"table {family} {table} {{\n" f"table {family} {table} {{",
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n" f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}",
f"}}\n" f"}}",
) ]
tmp = None
try: try:
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_chain_create_", suffix=".nft") as tf: for ln in script_lines:
tmp = tf.name nft_run_or_raise(ln)
tf.write(script) logger.debug("created chain %s in %s.%s via fallback lines", chain, family, table)
tf.flush() except RuntimeError as e2:
os.fsync(tf.fileno()) logger.error("fallback for creating chain failed: %s (original: %s)", e2, e_chain)
nft_run_or_raise(f"-f {tmp}") raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2
finally:
if tmp and os.path.exists(tmp): logger.info("table/chain ensured: %s.%s/%s", family, table, chain)
try:
os.remove(tmp)
except Exception:
pass
def _extract_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
# ---------------------- reconstruct rules from nft JSON (best-effort) ----------------------
def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]: def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
""" """
Best-effort mapping from nft expression json to our RuleModel-like dict. Heuristic mapping of nft expression JSON to our RuleModel fields.
This is heuristic: nft JSON shapes differ across kernel/libnftables versions. Covers common shapes: meta (iif/oif/length), cmp/payload (ip proto and ports), verdict (action).
We cover common patterns: meta (iif/oif/length), payload/cmp for ip proto and ports, verdict for action. This is best-effort; complex expressions may not map perfectly.
""" """
match: Dict[str, Any] = {} match: Dict[str, Any] = {}
action: Dict[str, Any] = {} action: Dict[str, Any] = {}
# iterate expressions; keep simple heuristics
for e in exprs: for e in exprs:
if "meta" in e: if "meta" in e:
m = e["meta"] m = e["meta"]
# keys differ; check for common names
key = m.get("key") or m.get("type") key = m.get("key") or m.get("type")
v = m.get("v") or m.get("s") or m.get("value")
if key in ("iifname", "iif"): if key in ("iifname", "iif"):
v = m.get("v") or m.get("s") or m.get("value")
if v: if v:
match["iif"] = v match["iif"] = v
elif key in ("oifname", "oif"): elif key in ("oifname", "oif"):
v = m.get("v") or m.get("s") or m.get("value")
if v: if v:
match["oif"] = v match["oif"] = v
elif key == "length": elif key == "length":
v = m.get("v") or m.get("s") or m.get("value")
if v is not None: if v is not None:
match["meta_length"] = v match["meta_length"] = v
elif "payload" in e: elif "payload" in e:
# payload indicates reading bytes of header; usually followed by a 'cmp' comparing to immediate # payload describes read of header bytes; often next 'cmp' compares it
# store payload description to use when we see a cmp # we tag the payload on the expression to help cmp heuristics
# flatten payload into a marker for later cmp detection e["_payload_hint"] = e["payload"]
e_payload = e["payload"]
e["_seen_payload"] = e_payload
elif "cmp" in e: elif "cmp" in e:
cmp = e["cmp"] cmp = e["cmp"]
# cmp can have 'left'/'right' or 'data' fields
left = cmp.get("left") left = cmp.get("left")
right = cmp.get("right") right = cmp.get("right")
# helper to pull immediate numeric value
def _extract_immediate(node): def _extract_immediate(node):
if not node: if not node or not isinstance(node, dict):
return None return None
if isinstance(node, dict): for k in ("immediate", "value", "data", "s", "v"):
for k in ("immediate", "value", "data", "s", "v"): if k in node:
if k in node: val = node[k]
val = node[k] if isinstance(val, str) and val.startswith("0x"):
# hex string -> int try:
if isinstance(val, str) and val.startswith("0x"): return int(val, 16)
try: except Exception:
return int(val, 16) return val
except Exception: return val
return val
return val
return None return None
imm_left = _extract_immediate(left) imm_left = _extract_immediate(left)
imm_right = _extract_immediate(right) imm_right = _extract_immediate(right)
# ip proto numeric likely in 1..255
# If either immediate looks like small numeric, assume ip_proto
for imm in (imm_left, imm_right): for imm in (imm_left, imm_right):
if isinstance(imm, int) and 0 < imm < 256: if isinstance(imm, int) and 0 < imm < 256:
# set ip_proto numeric
match["ip_proto"] = imm match["ip_proto"] = imm
break break
# port heuristics (1..65535)
# If immediate value looks like a TCP/UDP port (typical 1..65535) and payload context indicates tcp/udp dport,
# it's hard to be 100% sure; we use heuristic: if cmp mentions 'dport' or payload had offset consistent with port,
# or imm is in port range but no ip_proto set, we attempt to set tcp_dport/udp_dport.
imm = imm_left if imm_left is not None else imm_right imm = imm_left if imm_left is not None else imm_right
if isinstance(imm, int) and 0 < imm <= 65535: if isinstance(imm, int) and 0 < imm <= 65535:
# if we previously saw payload indicating tcp or udp, detect from payload description # best-effort assign to tcp_dport (common case)
# naive heuristic: if any seen payload mentions 'tcp' or 'udp' in its dict then assign accordingly # more advanced heuristics could inspect payload hints
# otherwise assign tcp_dport by default (best-effort) match.setdefault("tcp_dport", imm)
assigned = False
for ev in (left, right):
if isinstance(ev, dict):
# look for hints
if "payload" in ev:
pd = ev["payload"]
if isinstance(pd, dict) and ("tcp" in str(pd).lower()):
match["tcp_dport"] = imm
assigned = True
break
if not assigned:
# fallback to tcp_dport heuristic
match.setdefault("tcp_dport", imm)
elif "verdict" in e: elif "verdict" in e:
v = e["verdict"] v = e["verdict"]
# typical shapes: {"verdict":"accept"} or {"verdict":{"type":"drop"}} # shapes vary: dict or string
if isinstance(v, dict): if isinstance(v, dict):
t = v.get("type") or v.get("kind") t = v.get("type") or v.get("kind")
if t: if t:
action["type"] = t action["type"] = t
# redirect/queue handling may vary; try to extract # redirect / queue extra fields vary
if "to" in v: if "to" in v:
action["type"] = "redirect" action["type"] = "redirect"
action["redirect_port"] = v.get("to") action["redirect_port"] = v.get("to")
@@ -268,25 +227,18 @@ def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
action["queue_num"] = v.get("queue") action["queue_num"] = v.get("queue")
elif isinstance(v, str): elif isinstance(v, str):
action["type"] = v action["type"] = v
elif "jump" in e: # other expression types intentionally ignored for stateless reconstruction
# jump is effectively a control flow; not modeled
pass
elif "immediate" in e:
# sometimes immediate verdicts
pass
# there are many other expression types; above covers common ones
# Default action if none found
if "type" not in action: if "type" not in action:
action["type"] = "accept" # kernel often has policy accept if not specified # if kernel default, we assume accept
action["type"] = "accept"
return {"match": match, "action": action} return {"match": match, "action": action}
def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]: def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
""" """
Reconstruct a list of rules from the nft kernel listing (JSON). Reconstruct an ordered list of rule dicts from `nft list table` JSON via the binding.
Returns ordered list of dicts each matching RuleModel shape (without id). Returns list of dicts shaped like RuleModel (without id).
""" """
try: try:
out = nft_run_or_raise(f"list table {family} {table}") out = nft_run_or_raise(f"list table {family} {table}")
@@ -294,35 +246,33 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
logger.debug("no table %s.%s found when listing rules", family, table) logger.debug("no table %s.%s found when listing rules", family, table)
return [] return []
# libnftables returns spammy text sometimes; prefer JSON output # attempt to parse JSON output (binding can return textual JSON)
# attempt to parse as JSON if lib produced JSON; otherwise fallback to textual parsing
try: try:
# if out is JSON text produced by libnftables, it is already JSON representation of ruleset
data = json.loads(out) data = json.loads(out)
except Exception: except Exception:
# fallback: run with JSON mode by using the binding directly to request JSON output # force JSON mode in binding if previous parsing failed
_ensure_binding() _ensure_binding_available()
nft = Nftables() nft = Nftables()
nft.set_json_output(True) nft.set_json_output(True)
rc, out_json, err = nft.cmd(f"list table {family} {table}") rc, out_json, err = nft.cmd(f"list table {family} {table}")
if rc != 0: if rc != 0:
logger.debug("nft JSON list failed: %s", err) logger.debug("nft JSON listing failed: %s", err)
return [] return []
data = json.loads(out_json) data = json.loads(out_json)
results: List[Dict[str, Any]] = [] results: List[Dict[str, Any]] = []
# nft JSON structure: {"nftables":[ { "table":...}, { "chain":...}, { "rule": {...} }, ... ]}
for item in data.get("nftables", []): for item in data.get("nftables", []):
if "rule" not in item: if "rule" not in item:
continue continue
rule_obj = item["rule"] rule_obj = item["rule"]
exprs = rule_obj.get("expr", []) or rule_obj.get("expr", []) # different bindings key names exprs = rule_obj.get("expr", []) or rule_obj.get("expressions", []) or []
recon = _reconstruct_rule_from_exprs(exprs) recon = _reconstruct_rule_from_exprs(exprs)
# create a RuleModel-like dict # the chain name may be in the rule metadata
chain_name = rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN
rule_dict = { rule_dict = {
"family": family, "family": family,
"table": table, "table": table,
"chain": rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN, "chain": chain_name,
"match": recon.get("match", {}), "match": recon.get("match", {}),
"action": recon.get("action", {}), "action": recon.get("action", {}),
} }
@@ -335,11 +285,10 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
@router.get("/rules") @router.get("/rules")
def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): 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) logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", family, table, chain)
# ensure chain exists (create if missing) to provide consistent output
try: try:
ensure_table_chain(family, table, chain) ensure_table_chain(family, table, chain)
except RuntimeError as e: except RuntimeError as e:
logger.error("failed to ensure table/chain: %s", 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}") raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {e}")
rules = list_rules_from_nft(family, table) rules = list_rules_from_nft(family, table)
@@ -356,15 +305,15 @@ def put_rules(
chain: Optional[str] = DEFAULT_CHAIN, chain: Optional[str] = DEFAULT_CHAIN,
): ):
""" """
Replace entire ordered rule set. This implementation builds nft commands (without comments) Replace entire ordered ruleset. Validates provided rules, then builds nft commands
from the provided rules and applies via 'nft -f' using libnftables fallback. and applies them line-by-line via libnftables binding (no '-f' token passed into binding).
""" """
# optimistic concurrency # optimistic concurrency
global _current_version global _current_version
if if_match is not None and _current_version is not None and if_match != _current_version: 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") raise HTTPException(status_code=409, detail="version mismatch; fetch latest rules and retry")
# validate and ensure table/chain exist # validate per-rule family/table/chain (quick checks)
for r in rules: for r in rules:
if r.family and r.family != family: if r.family and r.family != family:
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}") raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}")
@@ -373,21 +322,24 @@ def put_rules(
if r.chain and r.chain != chain: if r.chain and r.chain != chain:
raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}") raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}")
# ensure table/chain exist first
try: try:
ensure_table_chain(family, table, chain) ensure_table_chain(family, table, chain)
except RuntimeError as e: except RuntimeError as e:
logger.error("failed to ensure table/chain: %s", e) logger.error("failed to ensure table/chain: %s", e)
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))
# Build nft script lines (no comments) # build script lines (flush + ordered add rules)
script_lines: List[str] = [] script_lines: List[str] = []
script_lines.append(f"flush chain {family} {table} {chain}") script_lines.append(f"flush chain {family} {table} {chain}")
for r in rules: for r in rules:
# ensure id exists only for client convenience; not stored # assign client-side id if missing (not stored)
if not r.id: if not r.id:
r.id = str(uuid.uuid4()) r.id = str(uuid.uuid4())
# build add rule command using the same helpers
match_frag = [] # build match fragments
match_frag: List[str] = []
if r.match.iif: if r.match.iif:
match_frag += ["iif", f'"{r.match.iif}"'] match_frag += ["iif", f'"{r.match.iif}"']
if r.match.oif: if r.match.oif:
@@ -401,43 +353,48 @@ def put_rules(
if r.match.udp_dport: if r.match.udp_dport:
match_frag += ["udp", "dport", str(r.match.udp_dport)] match_frag += ["udp", "dport", str(r.match.udp_dport)]
action_frag = [] # build action fragment
if r.action.type == "drop": action_frag: List[str] = []
a_type = r.action.type or "accept"
if a_type == "drop":
action_frag = ["drop"] action_frag = ["drop"]
elif r.action.type == "accept" or not r.action.type: elif a_type == "accept":
action_frag = ["accept"] action_frag = ["accept"]
elif r.action.type == "queue": elif a_type == "queue":
num = r.action.queue_num if r.action.queue_num is not None else 0 num = r.action.queue_num if r.action.queue_num is not None else 0
action_frag = ["queue", "num", str(num)] action_frag = ["queue", "num", str(num)]
elif r.action.type == "redirect": elif a_type == "redirect":
if r.action.redirect_port is None: if r.action.redirect_port is None:
raise HTTPException(status_code=400, detail="redirect action requires redirect_port") raise HTTPException(status_code=400, detail="redirect action requires redirect_port")
action_frag = ["redirect", "to", f":{r.action.redirect_port}"] action_frag = ["redirect", "to", f":{r.action.redirect_port}"]
else: else:
raise HTTPException(status_code=400, detail=f"unsupported action type: {r.action.type}") raise HTTPException(status_code=400, detail=f"unsupported action type: {a_type}")
parts = ["add", "rule", r.family, r.table, r.chain] parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag
parts += match_frag
parts += action_frag
script_lines.append(" ".join(parts)) script_lines.append(" ".join(parts))
script = "\n".join(script_lines) + "\n" # apply script lines via binding (line-by-line)
tmpfile_path: Optional[str] = None tmpfile_path: Optional[str] = None
try: 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: with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf:
tmpfile_path = tf.name tmpfile_path = tf.name
tf.write(script) tf.write("\n".join(script_lines) + "\n")
tf.flush() tf.flush()
os.fsync(tf.fileno()) os.fsync(tf.fileno())
logger.info("wrote nft script to %s; applying...", tmpfile_path) logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path)
try: try:
# use libnftables binding to apply file if possible, otherwise lib will call nft -f underneath for ln in script_lines:
nft_run_or_raise(f"-f {tmpfile_path}") ln = ln.strip()
if not ln:
continue
nft_run_or_raise(ln)
except RuntimeError as e: except RuntimeError as e:
logger.error("failed applying nft script: %s", 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}") raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}")
# success -> bump version
_current_version = str(uuid.uuid4()) _current_version = str(uuid.uuid4())
logger.info("applied nft ruleset successfully; version=%s", _current_version) logger.info("applied nft ruleset successfully; version=%s", _current_version)
return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules)) return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))