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.
Only uses kernel-stored info (nft) as the source of truth.
Stateless FastAPI nftables router that uses libnftables binding.
Endpoints (mounted under /nft):
- GET /rules -> reconstruct rule objects from nft kernel state (best-effort)
- PUT /rules -> replace entire ordered ruleset (applies via nft -f)
- 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)
No persistence, no comments used for mapping.
"""
from fastapi import APIRouter, HTTPException, Header, Request
from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any, Union, Tuple
from typing import Optional, List, Dict, Any
import uuid
import json
import tempfile
import os
import logging
# libnftables binding
# try import libnftables binding
try:
from nftables import Nftables
except Exception:
Nftables = None # will raise when used
Nftables = None # clearer error raised when attempting to use binding
router = APIRouter(prefix="/nft", tags=["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_CHAIN = "forward"
DEFAULT_FAMILY = "bridge"
_current_version: Optional[str] = None
# ---------------------- Pydantic models ----------------------
class MatchModel(BaseModel):
iif: Optional[str] = None
oif: Optional[str] = None
meta_length: Optional[Union[int, str]] = None
ip_proto: Optional[Union[int, str]] = None
meta_length: Optional[Any] = None # int or range string like "100-200"
ip_proto: Optional[Any] = None # numeric or name
tcp_dport: Optional[int] = None
udp_dport: Optional[int] = None
@@ -52,8 +47,8 @@ class ActionModel(BaseModel):
class RuleModel(BaseModel):
# note: since we don't persist IDs, id remains optional
id: Optional[str] = Field(None, description="optional client-provided id (not stored by server)")
# 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)
table: Optional[str] = Field(DEFAULT_TABLE)
chain: Optional[str] = Field(DEFAULT_CHAIN)
@@ -68,198 +63,162 @@ class ReplaceResult(BaseModel):
# ---------------------- nft binding helpers ----------------------
def _ensure_binding():
def _ensure_binding_available():
if Nftables is None:
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]:
"""Run a libnftables command and return (rc, stdout, stderr)."""
_ensure_binding()
def nft_cmd(cmd: str):
"""
Run a libnftables command string.
Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error.
"""
_ensure_binding_available()
nft = Nftables()
try:
rc, out, err = nft.cmd(cmd)
# ensure strings
# 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("nft binding call failed for command: %s", cmd)
raise RuntimeError(f"nft binding failed: {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:
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
# ---------------------- reconstruction logic ----------------------
# ---------------------- chain/table ensure / reconstruction ----------------------
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
if direct add fails. Raises RuntimeError on failure.
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; trying -f fallback", e_table)
script = f"table {family} {table} {{ }}\n"
tmp = None
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:
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_tbl_create_", suffix=".nft") as tf:
tmp = tf.name
tf.write(script)
tf.flush()
os.fsync(tf.fileno())
nft_run_or_raise(f"-f {tmp}")
finally:
if tmp and os.path.exists(tmp):
try:
os.remove(tmp)
except Exception:
pass
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
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:
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; trying -f fallback", e_chain)
script = (
f"table {family} {table} {{\n"
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n"
f"}}\n"
)
tmp = None
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:
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_chain_create_", suffix=".nft") as tf:
tmp = tf.name
tf.write(script)
tf.flush()
os.fsync(tf.fileno())
nft_run_or_raise(f"-f {tmp}")
finally:
if tmp and os.path.exists(tmp):
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
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) ----------------------
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.
This is heuristic: nft JSON shapes differ across kernel/libnftables versions.
We cover common patterns: meta (iif/oif/length), payload/cmp for ip proto and ports, verdict for action.
Heuristic mapping of nft expression JSON to our RuleModel fields.
Covers common shapes: meta (iif/oif/length), cmp/payload (ip proto and ports), verdict (action).
This is best-effort; complex expressions may not map perfectly.
"""
match: Dict[str, Any] = {}
action: Dict[str, Any] = {}
# iterate expressions; keep simple heuristics
for e in exprs:
if "meta" in e:
m = e["meta"]
# keys differ; check for common names
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"):
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
# payload describes read of header bytes; often next 'cmp' compares it
# we tag the payload on the expression to help cmp heuristics
e["_payload_hint"] = 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:
if not node or not isinstance(node, dict):
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:
return int(val, 16)
except Exception:
return val
return val
for k in ("immediate", "value", "data", "s", "v"):
if k in node:
val = node[k]
if isinstance(val, str) and val.startswith("0x"):
try:
return int(val, 16)
except Exception:
return val
return val
return None
imm_left = _extract_immediate(left)
imm_right = _extract_immediate(right)
# If either immediate looks like small numeric, assume ip_proto
# ip proto numeric likely in 1..255
for imm in (imm_left, imm_right):
if isinstance(imm, int) and 0 < imm < 256:
# set ip_proto numeric
match["ip_proto"] = imm
break
# 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.
# port heuristics (1..65535)
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)
# best-effort assign to tcp_dport (common case)
# more advanced heuristics could inspect payload hints
match.setdefault("tcp_dport", imm)
elif "verdict" in e:
v = e["verdict"]
# typical shapes: {"verdict":"accept"} or {"verdict":{"type":"drop"}}
# shapes vary: dict or string
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
# redirect / queue extra fields vary
if "to" in v:
action["type"] = "redirect"
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")
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
# other expression types intentionally ignored for stateless reconstruction
# Default action if none found
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}
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).
Reconstruct an ordered list of rule dicts from `nft list table` JSON via the binding.
Returns list of dicts shaped like RuleModel (without id).
"""
try:
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)
return []
# libnftables returns spammy text sometimes; prefer JSON output
# attempt to parse as JSON if lib produced JSON; otherwise fallback to textual parsing
# attempt to parse JSON output (binding can return textual JSON)
try:
# if out is JSON text produced by libnftables, it is already JSON representation of ruleset
data = json.loads(out)
except Exception:
# fallback: run with JSON mode by using the binding directly to request JSON output
_ensure_binding()
# force JSON mode in binding if previous parsing failed
_ensure_binding_available()
nft = Nftables()
nft.set_json_output(True)
rc, out_json, err = nft.cmd(f"list table {family} {table}")
if rc != 0:
logger.debug("nft JSON list failed: %s", err)
logger.debug("nft JSON listing failed: %s", err)
return []
data = json.loads(out_json)
results: List[Dict[str, Any]] = []
# nft JSON structure: {"nftables":[ { "table":...}, { "chain":...}, { "rule": {...} }, ... ]}
for item in data.get("nftables", []):
if "rule" not in item:
continue
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)
# 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 = {
"family": family,
"table": table,
"chain": rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN,
"chain": chain_name,
"match": recon.get("match", {}),
"action": recon.get("action", {}),
}
@@ -335,11 +285,10 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
@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)
# ensure chain exists (create if missing) to provide consistent output
try:
ensure_table_chain(family, table, chain)
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}")
rules = list_rules_from_nft(family, table)
@@ -356,15 +305,15 @@ def put_rules(
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.
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
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
# validate per-rule family/table/chain (quick checks)
for r in rules:
if r.family and 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:
raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}")
# ensure table/chain exist first
try:
ensure_table_chain(family, table, chain)
except RuntimeError as e:
logger.error("failed to ensure table/chain: %s", 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.append(f"flush chain {family} {table} {chain}")
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:
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:
match_frag += ["iif", f'"{r.match.iif}"']
if r.match.oif:
@@ -401,43 +353,48 @@ def put_rules(
if r.match.udp_dport:
match_frag += ["udp", "dport", str(r.match.udp_dport)]
action_frag = []
if r.action.type == "drop":
# build action fragment
action_frag: List[str] = []
a_type = r.action.type or "accept"
if a_type == "drop":
action_frag = ["drop"]
elif r.action.type == "accept" or not r.action.type:
elif a_type == "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
action_frag = ["queue", "num", str(num)]
elif r.action.type == "redirect":
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: {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 += match_frag
parts += action_frag
parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag
script_lines.append(" ".join(parts))
script = "\n".join(script_lines) + "\n"
# apply script lines via binding (line-by-line)
tmpfile_path: Optional[str] = None
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(script)
tf.write("\n".join(script_lines) + "\n")
tf.flush()
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:
# use libnftables binding to apply file if possible, otherwise lib will call nft -f underneath
nft_run_or_raise(f"-f {tmpfile_path}")
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: %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}")
# 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))