add libnft, simpler api
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s

This commit is contained in:
2026-01-10 19:38:55 +01:00
parent cc61aeade6
commit 09d64e6eec
2 changed files with 285 additions and 257 deletions

Binary file not shown.

View File

@@ -1,106 +1,116 @@
# fastapi_nft_replace.py # fastapi_nft_stateless.py
""" """
nftables router for FastAPI to manage nftables bridge rules dynamically. Stateless FastAPI nftables router.
Only uses kernel-stored info (nft) as the source of truth.
This version will attempt to create any kernel/runtime prerequisites:
- load 'bridge' and 'br_netfilter' modules (via modprobe)
- enable sysctls net.bridge.bridge-nf-call-iptables and net.bridge.bridge-nf-call-ip6tables
- create the nft table and chain (with nft add ... and nft -f fallback)
Endpoints (mounted under /nft): Endpoints (mounted under /nft):
- GET /rules -> list active rules (read from nftables) - GET /rules -> reconstruct rule objects from nft kernel state (best-effort)
- POST /rules -> add one or many rules (append) - PUT /rules -> replace entire ordered ruleset (applies via nft -f)
- DELETE /rules -> delete one or many rules by id
- PUT /rules -> replace entire ordered rule set via nft -f (returns new version) No persistence, no comments used for mapping.
""" """
from fastapi import APIRouter, HTTPException 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, Union, Tuple
import subprocess
import uuid import uuid
import json import json
import base64
import re
import tempfile import tempfile
import os import os
import logging import logging
# Router and logger --------------------------------------------------------- # libnftables binding
router = APIRouter() try:
from nftables import Nftables
except Exception:
Nftables = None # will raise when used
router = APIRouter(prefix="/nft", tags=["nftables"])
logger = logging.getLogger("nftables") logger = logging.getLogger("nftables")
logger.debug("nftables router module loaded") logger.debug("nftables stateless router loaded")
# Defaults and in-memory version token -------------------------------------- # Defaults
DEFAULT_TABLE = "mitm_tbl" DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward" DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge" DEFAULT_FAMILY = "bridge"
# in-memory version token updated on successful PUT
_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 # allow integers or ranges like "100-200" meta_length: Optional[Union[int, str]] = None
ip_proto: Optional[str] = None ip_proto: Optional[Union[int, str]] = None
tcp_dport: Optional[int] = None tcp_dport: Optional[int] = None
udp_dport: Optional[int] = None udp_dport: Optional[int] = None
class ActionModel(BaseModel): class ActionModel(BaseModel):
type: str # drop | accept | queue | redirect type: Optional[str] = None # drop | accept | queue | redirect
queue_num: Optional[int] = None queue_num: Optional[int] = None
redirect_port: Optional[int] = None redirect_port: Optional[int] = None
class RuleModel(BaseModel): class RuleModel(BaseModel):
id: Optional[str] = Field(None, description="optional rule id; generated if missing") # note: since we don't persist IDs, id remains optional
id: Optional[str] = Field(None, description="optional client-provided id (not stored by server)")
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)
match: MatchModel match: MatchModel
action: ActionModel action: ActionModel
class ReplaceResult(BaseModel): class ReplaceResult(BaseModel):
version: str version: str
applied: bool applied: bool
rules_count: int rules_count: int
# ---------------------- Utilities ---------------------- # ---------------------- nft binding helpers ----------------------
def run_nft(args: List[str]) -> Tuple[str, str]: def _ensure_binding():
""" if Nftables is None:
Run nft with given args. raise RuntimeError(
Returns (stdout, stderr). "python nftables binding not available. Install python3-nftables (system package) or pip-nftables."
Raises RuntimeError on non-zero exit with stderr included. )
"""
logger.debug("running nft: %s", " ".join(["nft"] + args))
try:
proc = subprocess.run(["nft"] + args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True)
logger.debug("nft stdout: %s", proc.stdout.strip())
return proc.stdout, proc.stderr
except subprocess.CalledProcessError as e:
logger.error("nft failed: %s -- %s", " ".join(e.cmd), e.stderr.strip())
raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}")
def nft_cmd(cmd: str) -> Tuple[int, str, str]:
"""Run a libnftables command and return (rc, stdout, stderr)."""
_ensure_binding()
nft = Nftables()
try:
rc, out, err = nft.cmd(cmd)
# ensure 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("nft binding call failed for command: %s", cmd)
raise RuntimeError(f"nft binding failed: {e}")
def nft_run_or_raise(cmd: str) -> str:
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:
raise RuntimeError(err or f"nft command {cmd} failed with rc={rc}")
return out
# ---------------------- reconstruction logic ----------------------
def ensure_table_chain(family: str, table: str, chain: str) -> None: def ensure_table_chain(family: str, table: str, chain: str) -> None:
""" """
Ensure the nft table and chain exist. On serious failures this raises RuntimeError. Ensure table and chain exist. Use libnftables commands; fallback to nft -f temp file
This function: if direct add fails. Raises RuntimeError on failure.
- attempts `nft add table` and `nft add chain`
- if those fail, tries to apply a tiny nft script with `nft -f` to create table and chain
""" """
logger.info("ensuring table %s.%s exists", family, table) logger.info("ensuring table %s.%s exists", family, table)
# Try simple add table first
try: try:
run_nft(["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("nft add table failed for %s.%s; trying nft -f fallback: %s", family, table, e_table) logger.info("add table failed: %s; trying -f fallback", e_table)
# fallback create via nft -f script
script = f"table {family} {table} {{ }}\n" script = f"table {family} {table} {{ }}\n"
tmp = None tmp = None
try: try:
@@ -109,12 +119,7 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
tf.write(script) tf.write(script)
tf.flush() tf.flush()
os.fsync(tf.fileno()) os.fsync(tf.fileno())
try: nft_run_or_raise(f"-f {tmp}")
run_nft(["-f", tmp])
logger.debug("created table %s.%s via nft -f", family, table)
except RuntimeError as e2:
logger.error("nft -f fallback to create table failed: %s (original: %s)", e2, e_table)
raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2
finally: finally:
if tmp and os.path.exists(tmp): if tmp and os.path.exists(tmp):
try: try:
@@ -122,13 +127,13 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
except Exception: except Exception:
pass pass
# Now ensure chain exists logger.info("ensuring chain %s in table %s", chain, table)
logger.info("ensuring chain %s in table %s exists", chain, table)
try: try:
run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"]) nft_run_or_raise(
logger.debug("created chain %s in %s.%s via add chain", chain, family, table) f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}'
)
except RuntimeError as e_chain: except RuntimeError as e_chain:
logger.info("nft add chain failed for %s in %s.%s; trying nft -f fallback: %s", chain, family, table, e_chain) logger.info("add chain failed: %s; trying -f fallback", e_chain)
script = ( script = (
f"table {family} {table} {{\n" f"table {family} {table} {{\n"
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n" f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n"
@@ -141,12 +146,7 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
tf.write(script) tf.write(script)
tf.flush() tf.flush()
os.fsync(tf.fileno()) os.fsync(tf.fileno())
try: nft_run_or_raise(f"-f {tmp}")
run_nft(["-f", tmp])
logger.debug("created chain %s in %s.%s via nft -f", chain, family, table)
except RuntimeError as e2:
logger.error("nft -f fallback to create chain failed: %s (original: %s)", e2, e_chain)
raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2
finally: finally:
if tmp and os.path.exists(tmp): if tmp and os.path.exists(tmp):
try: try:
@@ -154,210 +154,217 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
except Exception: except Exception:
pass pass
logger.info("table/chain ensured: %s.%s/%s", family, table, chain)
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
def encode_rule_comment(rule: Dict[str, Any]) -> str: def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
# store rule JSON as base64 to avoid quoting/escaping issues inside nft comment """
j = json.dumps(rule, separators=(",", ":")) Best-effort mapping from nft expression json to our RuleModel-like dict.
b = base64.b64encode(j.encode()).decode() This is heuristic: nft JSON shapes differ across kernel/libnftables versions.
rid = rule.get("id") or "" We cover common patterns: meta (iif/oif/length), payload/cmp for ip proto and ports, verdict for action.
return f"mitm_id:{rid} mitm_json:{b}" """
match: Dict[str, Any] = {}
action: Dict[str, Any] = {}
# iterate expressions; keep simple heuristics
def decode_comment_payload(comment: str) -> Optional[Dict[str, Any]]: for e in exprs:
# expects comment like: mitm_id:<id> mitm_json:<base64> if "meta" in e:
m = e["meta"]
# keys differ; check for common names
key = m.get("key") or m.get("type")
if key in ("iifname", "iif"):
v = m.get("v") or m.get("s") or m.get("value")
if v:
match["iif"] = v
elif key in ("oifname", "oif"):
v = m.get("v") or m.get("s") or m.get("value")
if v:
match["oif"] = v
elif key == "length":
v = m.get("v") or m.get("s") or m.get("value")
if v is not None:
match["meta_length"] = v
elif "payload" in e:
# payload indicates reading bytes of header; usually followed by a 'cmp' comparing to immediate
# store payload description to use when we see a cmp
# flatten payload into a marker for later cmp detection
e_payload = e["payload"]
e["_seen_payload"] = e_payload
elif "cmp" in e:
cmp = e["cmp"]
# cmp can have 'left'/'right' or 'data' fields
left = cmp.get("left")
right = cmp.get("right")
# helper to pull immediate numeric value
def _extract_immediate(node):
if not node:
return None
if isinstance(node, dict):
for k in ("immediate", "value", "data", "s", "v"):
if k in node:
val = node[k]
# hex string -> int
if isinstance(val, str) and val.startswith("0x"):
try: try:
parts = comment.split() return int(val, 16)
kv = {p.split(":", 1)[0]: p.split(":", 1)[1] for p in parts if ":" in p}
b64 = kv.get("mitm_json")
if not b64:
return None
j = base64.b64decode(b64.encode()).decode()
return json.loads(j)
except Exception: except Exception:
logger.debug("failed to decode comment payload: %s", comment) return val
return val
return None return None
imm_left = _extract_immediate(left)
imm_right = _extract_immediate(right)
def build_nft_match_fragment(match: MatchModel) -> List[str]: # If either immediate looks like small numeric, assume ip_proto
frag: List[str] = [] for imm in (imm_left, imm_right):
if match.iif: if isinstance(imm, int) and 0 < imm < 256:
frag += ["iif", f'"{match.iif}"'] # set ip_proto numeric
if match.oif: match["ip_proto"] = imm
frag += ["oif", f'"{match.oif}"'] break
if match.meta_length is not None:
frag += ["meta", "length", str(match.meta_length)]
if match.ip_proto:
frag += ["ip", "protocol", 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
# 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
if isinstance(imm, int) and 0 < imm <= 65535:
# if we previously saw payload indicating tcp or udp, detect from payload description
# naive heuristic: if any seen payload mentions 'tcp' or 'udp' in its dict then assign accordingly
# otherwise assign tcp_dport by default (best-effort)
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:
v = e["verdict"]
# typical shapes: {"verdict":"accept"} or {"verdict":{"type":"drop"}}
if isinstance(v, dict):
t = v.get("type") or v.get("kind")
if t:
action["type"] = t
# redirect/queue handling may vary; try to extract
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
elif "jump" in e:
# 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
def build_nft_action_fragment(action: ActionModel) -> List[str]: # Default action if none found
if action.type == "drop": if "type" not in action:
return ["drop"] action["type"] = "accept" # kernel often has policy accept if not specified
if action.type == "accept":
return ["accept"]
if action.type == "queue":
num = action.queue_num if action.queue_num is not None else 0
return ["queue", "num", str(num)]
if action.type == "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}")
return {"match": match, "action": action}
def nft_rule_line_from_model(rule: RuleModel) -> str:
"""
Produce a single-line nft command:
add rule <family> <table> <chain> <match...> <action...> comment "<encoded>"
"""
match_frag = build_nft_match_fragment(rule.match)
action_frag = build_nft_action_fragment(rule.action)
comment = encode_rule_comment(rule.dict())
parts = ["add", "rule", rule.family, rule.table, rule.chain]
parts += match_frag
parts += action_frag
return " ".join(parts) + f' comment "{comment}"'
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).
Returns ordered list of dicts each matching RuleModel shape (without id).
"""
try: try:
out, _ = run_nft(["list", "table", family, table]) out = nft_run_or_raise(f"list table {family} {table}")
except RuntimeError: except RuntimeError:
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 []
results: List[Dict[str, Any]] = [] # libnftables returns spammy text sometimes; prefer JSON output
for line in out.splitlines(): # attempt to parse as JSON if lib produced JSON; otherwise fallback to textual parsing
line = line.strip()
if "comment" in line and "mitm_json:" in line:
try: try:
first_quote = line.index('"') # if out is JSON text produced by libnftables, it is already JSON representation of ruleset
last_quote = line.rindex('"') data = json.loads(out)
comment_str = line[first_quote + 1:last_quote] except Exception:
except ValueError: # fallback: run with JSON mode by using the binding directly to request JSON output
comment_str = line.split("comment", 1)[1].strip() _ensure_binding()
nft = Nftables()
parsed = decode_comment_payload(comment_str) nft.set_json_output(True)
if parsed is not None: rc, out_json, err = nft.cmd(f"list table {family} {table}")
results.append(parsed) if rc != 0:
logger.debug("listed %d nft rules from %s.%s", len(results), family, table) logger.debug("nft JSON list failed: %s", err)
return results
def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]:
ensure_table_chain(rule.family, rule.table, rule.chain)
if not rule.id:
rule.id = str(uuid.uuid4())
cmd_text = nft_rule_line_from_model(rule)
try:
# run as discrete args to avoid shell quoting issues
run_nft(cmd_text.split())
logger.info("added nft rule id=%s family=%s table=%s chain=%s", rule.id, rule.family, rule.table, rule.chain)
return {"id": rule.id, "status": "added"}
except RuntimeError as e:
logger.error("failed to add rule id=%s: %s", rule.id, e)
raise HTTPException(status_code=500, detail=str(e))
def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]:
try:
out, _ = run_nft(["list", "chain", family, table, chain, "-a"])
except RuntimeError:
logger.debug("no chain %s in table %s.%s when attempting delete", chain, family, table)
return [] return []
data = json.loads(out_json)
deleted: List[str] = [] results: List[Dict[str, Any]] = []
for line in out.splitlines(): # nft JSON structure: {"nftables":[ { "table":...}, { "chain":...}, { "rule": {...} }, ... ]}
if "comment" in line and "mitm_id:" in line: for item in data.get("nftables", []):
try: if "rule" not in item:
q1 = line.index('"')
q2 = line.index('"', q1 + 1)
comment_str = line[q1 + 1:q2]
except ValueError:
comment_str = line.split("comment", 1)[1]
parsed = decode_comment_payload(comment_str)
if not parsed:
continue continue
rid = parsed.get("id") rule_obj = item["rule"]
if rid in ids: exprs = rule_obj.get("expr", []) or rule_obj.get("expr", []) # different bindings key names
m = re.search(r"handle\s+(\d+)", line) recon = _reconstruct_rule_from_exprs(exprs)
if not m: # create a RuleModel-like dict
parts = line.split() rule_dict = {
if "handle" in parts: "family": family,
hi = parts.index("handle") "table": table,
if hi + 1 < len(parts): "chain": rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN,
handle = parts[hi + 1] "match": recon.get("match", {}),
else: "action": recon.get("action", {}),
continue }
else: results.append(rule_dict)
continue logger.debug("reconstructed %d rules from %s.%s", len(results), family, table)
else: return results
handle = m.group(1)
try:
run_nft(["delete", "rule", family, table, chain, "handle", handle])
logger.info("deleted nft rule id=%s handle=%s", rid, handle)
deleted.append(rid)
except RuntimeError:
logger.error("failed to delete nft rule id=%s handle=%s", rid, handle)
continue
return deleted
# ---------------------- API endpoints ---------------------- # ---------------------- API endpoints ----------------------
@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):
"""
Ensure the table/chain exist (try to create them if missing), then return
the rules from nftables. If ensure_table_chain fails we return 500 with a clear message.
"""
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 required table/chain exist before listing rules
ensure_table_chain(family, table, chain) ensure_table_chain(family, table, chain)
except RuntimeError as e: except RuntimeError as e:
logger.error("failed to create/ensure nft table/chain %s.%s/%s: %s", family, table, chain, e) logger.error("failed to ensure table/chain: %s", e)
# return an HTTP 500 so the frontend knows setup failed
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}")
# now safe to list rules
rules = list_rules_from_nft(family, table) rules = list_rules_from_nft(family, table)
return {"count": len(rules), "rules": rules, "version": _current_version} return {"count": len(rules), "rules": rules, "version": _current_version}
@router.post("/rules")
def post_rules(payload: Union[RuleModel, List[RuleModel]]):
rules = payload if isinstance(payload, list) else [payload]
results = []
for r in rules:
res = add_rule_to_nft(r)
results.append(res)
return {"results": results}
@router.delete("/rules")
def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
ids = [payload] if isinstance(payload, str) else payload
deleted = delete_rules_by_ids(ids, family, table, chain)
results = [{"id": i, "deleted": i in deleted} for i in ids]
return {"results": results}
@router.put("/rules", response_model=ReplaceResult) @router.put("/rules", response_model=ReplaceResult)
def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY, def put_rules(
table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN): rules: List[RuleModel],
# validate per-rule family/table/chain if present request: Request,
if_match: Optional[str] = Header(None, alias="If-Match"),
family: Optional[str] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE,
chain: Optional[str] = DEFAULT_CHAIN,
):
"""
Replace entire ordered rule set. This implementation builds nft commands (without comments)
from the provided rules and applies via 'nft -f' using libnftables fallback.
"""
# optimistic concurrency
global _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")
# validate and ensure table/chain exist
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}")
@@ -366,32 +373,55 @@ def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY,
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 (this will attempt to modprobe + sysctl if needed)
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))
# ensure rule IDs # Build nft script lines (no comments)
script_lines: List[str] = []
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
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 = []
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 nft script lines action_frag = []
lines: List[str] = [] if r.action.type == "drop":
# flush chain (clear existing ordered rules) action_frag = ["drop"]
lines.append(f"flush chain {family} {table} {chain}") elif r.action.type == "accept" or not r.action.type:
action_frag = ["accept"]
elif r.action.type == "queue":
num = r.action.queue_num if r.action.queue_num is not None else 0
action_frag = ["queue", "num", str(num)]
elif r.action.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: {r.action.type}")
# add ordered rules parts = ["add", "rule", r.family, r.table, r.chain]
for r in rules: parts += match_frag
try: parts += action_frag
lines.append(nft_rule_line_from_model(r)) script_lines.append(" ".join(parts))
except Exception as e:
raise HTTPException(status_code=400, detail=f"invalid rule: {e}")
script = "\n".join(lines) + "\n"
script = "\n".join(script_lines) + "\n"
tmpfile_path: Optional[str] = None tmpfile_path: Optional[str] = None
try: try:
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:
@@ -401,15 +431,13 @@ def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY,
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...", tmpfile_path)
# apply the file
try: try:
run_nft(["-f", tmpfile_path]) # use libnftables binding to apply file if possible, otherwise lib will call nft -f underneath
nft_run_or_raise(f"-f {tmpfile_path}")
except RuntimeError as e: except RuntimeError as e:
logger.error("failed applying nft script: %s", e) logger.error("failed applying nft script: %s", 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 token
global _current_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))