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. Stateless FastAPI nftables router using pyroute2.nftables.
- 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: 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 fastapi import APIRouter, HTTPException, Header, Request
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any from typing import Optional, List, Dict, Any, Union
import uuid
import json
import tempfile
import os
import logging import logging
import json
import uuid
from enum import Enum
# try import libnftables binding # pyroute2 nftables
try: from pyroute2 import nftables
from nftables import Nftables
except Exception:
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 module loaded") logger.debug("nftables stateless-comment-enums router loaded")
# Defaults and version token # Defaults and version token
DEFAULT_TABLE = "mitm_tbl" DEFAULT_TABLE = "mitm_tbl"
@@ -30,30 +30,54 @@ DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge" DEFAULT_FAMILY = "bridge"
_current_version: Optional[str] = None _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 ---------------------- # ---------------------- 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[Any] = None # int or range string like "100-200" meta_length: Optional[Any] = None # int or range string
ip_proto: Optional[Any] = None # numeric or name # ip_proto may be Protocol enum, integer or arbitrary string (name)
ip_proto: Optional[Union[int, Protocol, 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: Optional[str] = None # drop | accept | queue | redirect type: ActionType # ActionType enum
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 is optional client-side convenience only (not persisted) id: Optional[str] = Field(None, description="optional client id; not persisted")
id: Optional[str] = Field(None, description="optional client id; not stored server-side") family: Optional[Family] = Field(Family.BRIDGE)
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
comment: Optional[str] = Field(None, description="optional human-readable comment stored in nft comment")
class ReplaceResult(BaseModel): class ReplaceResult(BaseModel):
@@ -62,216 +86,210 @@ class ReplaceResult(BaseModel):
rules_count: int rules_count: int
# ---------------------- nft binding helpers ---------------------- # ---------------------- nft wrapper ----------------------
def _ensure_binding_available(): class NFT:
if Nftables is None: def __init__(self):
raise RuntimeError( self.nft = nftables.NFTables()
"python nftables binding not available. Install system package `python3-nftables` or a compatible binding."
)
def run(self, cmd: str) -> Dict[str, Any]:
def nft_cmd(cmd: str): """
""" Run a single nft command string via pyroute2.NFTables.cmd.
Run a libnftables command string. Returns parsed JSON if possible, otherwise a dict with 'out' textual output.
Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error. Raises RuntimeError on failure.
""" """
_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} {{ }}"]
try: try:
for ln in script_lines: rc, out, err = self.nft.cmd(cmd)
nft_run_or_raise(ln) if isinstance(out, bytes):
logger.debug("created table %s.%s via fallback lines", family, table) out = out.decode(errors="ignore")
except RuntimeError as e2: if isinstance(err, bytes):
logger.error("fallback for creating table failed: %s (original: %s)", e2, e_table) err = err.decode(errors="ignore")
raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2 if rc != 0:
raise RuntimeError(err or f"nft cmd failed rc={rc}")
# try add chain if out:
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"):
try: try:
return int(val, 16) return json.loads(out)
except Exception: except Exception:
return val return {"out": out}
return val return {}
# sometimes immediate sits under {"immediate": {"value": ...}} except AttributeError:
if "right" in node or "left" in node: # Try json_cmd fallback if older pyroute2 version
# caller handles left/right structures out = self.nft.json_cmd(cmd)
return None 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 <family> <table> <chain>'.
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 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. Given a rule entry returned by NFTC.list_table, return a dict that resembles RuleModel
This function attempts to match many of the shapes produced by libnftables. (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] = {} match: Dict[str, Any] = {}
action: Dict[str, Any] = {} action: Dict[str, Any] = {}
# Track last payload hint (if present) so that subsequent cmp can be interpreted for e in exprs:
last_payload_hint: Optional[Dict[str, Any]] = None
for idx, e in enumerate(exprs):
# handle meta (iif/oif/length)
if "meta" in e: if "meta" in e:
m = e["meta"] m = e["meta"]
key = m.get("key") or m.get("type") or m.get("field") key = m.get("key") or m.get("type") or m.get("field")
v = m.get("v") or m.get("s") or m.get("value") v = m.get("v") or m.get("s") or m.get("value")
# sometimes v is dict
if isinstance(v, dict): if isinstance(v, dict):
v = _extract_immediate_value(v) or v.get("s") or v.get("v") v = v.get("value") or v.get("v") or v.get("s")
if key in ("iifname", "iif", "in", "iifname?"): if key in ("iifname", "iif", "in"):
if v: match["iif"] = v
match["iif"] = v
elif key in ("oifname", "oif", "out"): elif key in ("oifname", "oif", "out"):
if v: match["oif"] = v
match["oif"] = v
elif key in ("length", "len"): elif key in ("length", "len"):
if v is not None: match["meta_length"] = v
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
elif "cmp" in e or "match" in e: elif "cmp" in e or "match" in e:
cmp_obj = e.get("cmp") or e.get("match") or {} cmp_obj = e.get("cmp") or e.get("match") or {}
left = cmp_obj.get("left") left = cmp_obj.get("left")
right = cmp_obj.get("right") right = cmp_obj.get("right")
# extract immediate numeric if present def _extract_immediate(x):
imm_left = _extract_immediate_value(left) if not x or not isinstance(x, dict):
imm_right = _extract_immediate_value(right) return None
imm = imm_left if imm_left is not None else imm_right for k in ("immediate", "value", "data", "s", "v"):
if k in x:
# If either immediate is a small int treat as ip_proto val = x[k]
if isinstance(imm, int) and 0 < imm < 256: if isinstance(val, str) and val.startswith("0x"):
# prefer to store as number (frontend may show numeric) try:
match["ip_proto"] = imm return int(val, 16)
except Exception:
# If immediate looks like port (1-65535), try to detect target (tcp/udp) return val
if isinstance(imm, int) and 0 < imm <= 65535: return val
assigned = False return None
# heuristics: if left/right contains a payload hint referring to TCP/UDP or 'dport' strings: imm = _extract_immediate(left) or _extract_immediate(right)
for side in (left, right): if isinstance(imm, str):
if isinstance(side, dict): low = imm.lower()
# payload form used by libnftables can include 'protocol' or 'field' if low == "icmp":
pl = side.get("payload") or side.get("left", {}).get("payload") match["ip_proto"] = 1
if isinstance(pl, dict): elif low == "tcp":
protocol_hint = pl.get("protocol") or pl.get("proto") or pl.get("family") match["ip_proto"] = 6
field_hint = pl.get("field") or pl.get("meta") elif low == "udp":
sh = json.dumps(pl).lower() match["ip_proto"] = 17
if "tcp" in sh or "sport" in sh or "dport" in sh or "th" in sh: else:
match["tcp_dport"] = imm try:
assigned = True match["ip_proto"] = int(imm)
break except Exception:
if "udp" in sh or "udph" in sh or "udp." in sh: match["ip_proto"] = imm
match["udp_dport"] = imm if isinstance(imm, int):
assigned = True if 0 < imm < 256:
break match["ip_proto"] = imm
if not assigned: elif 0 < imm <= 65535:
# fallback: if ip_proto already indicates tcp(6) or udp(17), assign accordingly match.setdefault("tcp_dport", imm)
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
elif "verdict" in e or "immediate" in e or "return" in e: 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") v = e.get("verdict") or e.get("return") or e.get("immediate")
# normalize dict forms
if isinstance(v, dict): if isinstance(v, dict):
t = v.get("type") or v.get("kind") or v.get("verdict") t = v.get("type") or v.get("kind") or v.get("verdict")
if t: if t:
action["type"] = t action["type"] = t
# redirect/queue shaped differently across versions
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")
@@ -279,103 +297,128 @@ def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
action["type"] = "queue" action["type"] = "queue"
action["queue_num"] = v.get("queue") action["queue_num"] = v.get("queue")
elif isinstance(v, str): elif isinstance(v, str):
# strings like "accept" or "drop"
action["type"] = v 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: else:
# unknown expression type: log (DEBUG) for later tuning # ignore other expression types
logger.debug("unhandled nft expr type: %s", list(e.keys())) pass
# default action to accept if kernel default or none found
if "type" not in action: if "type" not in action:
action["type"] = "accept" action["type"] = "accept"
return {"match": match, "action": action} # convert action.type to ActionType enum if possible
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 []
try: 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: 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 [] 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]] = [] results: List[Dict[str, Any]] = []
for item in data.get("nftables", []): for item in entries:
if "rule" not in item: if not item:
continue continue
rule_obj = item["rule"] if "rule" in item:
# different versions might use 'expr', 'expressions' or 'expr' rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")}
exprs = rule_obj.get("expr") or rule_obj.get("exprs") or rule_obj.get("expressions") or [] else:
# if exprs is not a list but a dict (some shaped outputs) normalize rule_entry = item
if isinstance(exprs, dict): reconstructed = reconstruct_rule_from_rule_entry(rule_entry)
# sometimes expressions are nested inside a single dict; try to find array keys results.append(reconstructed)
for k in ("expr", "expressions", "exprs"):
v = exprs.get(k)
if isinstance(v, list):
exprs = v
break
else:
# give up and wrap
exprs = [exprs]
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 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 ---------------------- # ---------------------- 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[Union[str, Family]] = 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) 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: try:
ensure_table_chain(family, table, chain) ensure_table_chain(fam, table, chain)
except RuntimeError as e: except Exception as e:
logger.error("failed to ensure table/chain %s.%s/%s: %s", family, table, chain, e) logger.error("failed to ensure table/chain: %s", e)
raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {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} return {"count": len(rules), "rules": rules, "version": _current_version}
@@ -384,107 +427,46 @@ def put_rules(
rules: List[RuleModel], rules: List[RuleModel],
request: Request, request: Request,
if_match: Optional[str] = Header(None, alias="If-Match"), 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, table: Optional[str] = DEFAULT_TABLE,
chain: Optional[str] = DEFAULT_CHAIN, 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 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: 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 per-rule family/table/chain (quick checks) # validate per-rule family/table/chain
for r in rules: for r in rules:
if r.family and r.family != family: r_family_val = r.family.value if isinstance(r.family, Family) else r.family
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {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: if r.table and r.table != table:
raise HTTPException(status_code=400, detail=f"rule table mismatch: {r.table} != {table}") raise HTTPException(status_code=400, detail=f"rule table mismatch: {r.table} != {table}")
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 # ensure table/chain exist
try: try:
ensure_table_chain(family, table, chain) ensure_table_chain(fam, table, chain)
except RuntimeError as e: except Exception 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 script lines (flush + ordered add rules) # ensure rule ids for client convenience
script_lines: List[str] = []
script_lines.append(f"flush chain {family} {table} {chain}")
for r in rules: for r in rules:
# 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 match fragments # attempt to replace rules
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
try: try:
# keep a temp file copy for debugging if desired (optional) add_rules_replace_all(rules, fam, table, chain)
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf: except Exception as e:
tmpfile_path = tf.name logger.error("failed to apply rules: %s", e)
tf.write("\n".join(script_lines) + "\n") raise HTTPException(status_code=500, detail=str(e))
tf.flush()
os.fsync(tf.fileno())
logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path)
try: _current_version = str(uuid.uuid4())
for ln in script_lines: logger.info("applied nft ruleset successfully; version=%s", _current_version)
ln = ln.strip() return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))
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)