Files
mitm-webserver/backend/src/api/nftables_api.py
malmert 78b687087d
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s
improve
2026-01-10 20:06:09 +01:00

407 lines
16 KiB
Python

"""
Stateless FastAPI nftables router that uses libnftables binding.
- GET /nft/rules -> reconstruct rule objects from nft kernel state (best-effort)
- PUT /nft/rules -> replace entire ordered ruleset (applies via binding, line-by-line)
"""
from fastapi import APIRouter, HTTPException, Header, Request
from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any
import uuid
import json
import tempfile
import os
import logging
# try import libnftables binding
try:
from nftables import Nftables
except Exception:
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 module loaded")
# 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[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
class ActionModel(BaseModel):
type: Optional[str] = None # drop | accept | queue | redirect
queue_num: Optional[int] = None
redirect_port: Optional[int] = None
class RuleModel(BaseModel):
# id is optional client-side convenience only (not persisted)
id: Optional[str] = Field(None, description="optional client id; not stored server-side")
family: Optional[str] = Field(DEFAULT_FAMILY)
table: Optional[str] = Field(DEFAULT_TABLE)
chain: Optional[str] = Field(DEFAULT_CHAIN)
match: MatchModel
action: ActionModel
class ReplaceResult(BaseModel):
version: str
applied: bool
rules_count: int
# ---------------------- nft binding helpers ----------------------
def _ensure_binding_available():
if Nftables is None:
raise RuntimeError(
"python nftables binding not available. Install system package `python3-nftables` or a compatible binding."
)
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 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:
for ln in script_lines:
nft_run_or_raise(ln)
logger.debug("created table %s.%s via fallback lines", family, table)
except RuntimeError as e2:
logger.error("fallback for creating table failed: %s (original: %s)", e2, e_table)
raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2
# try add chain
logger.info("ensuring chain %s in table %s exists", chain, table)
try:
nft_run_or_raise(
f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}'
)
logger.debug("created chain %s in %s.%s via add chain", chain, family, table)
except RuntimeError as e_chain:
logger.info("add chain failed: %s. attempting fallback", e_chain)
script_lines = [
f"table {family} {table} {{",
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}",
f"}}",
]
try:
for ln in script_lines:
nft_run_or_raise(ln)
logger.debug("created chain %s in %s.%s via fallback lines", chain, family, table)
except RuntimeError as e2:
logger.error("fallback for creating chain failed: %s (original: %s)", e2, e_chain)
raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2
logger.info("table/chain ensured: %s.%s/%s", family, table, chain)
# ---------------------- reconstruct rules from nft JSON (best-effort) ----------------------
def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
"""
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] = {}
for e in exprs:
if "meta" in e:
m = e["meta"]
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 v:
match["iif"] = v
elif key in ("oifname", "oif"):
if v:
match["oif"] = v
elif key == "length":
if v is not None:
match["meta_length"] = v
elif "payload" in e:
# 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"]
left = cmp.get("left")
right = cmp.get("right")
def _extract_immediate(node):
if not node or not isinstance(node, dict):
return None
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)
# ip proto numeric likely in 1..255
for imm in (imm_left, imm_right):
if isinstance(imm, int) and 0 < imm < 256:
match["ip_proto"] = imm
break
# port heuristics (1..65535)
imm = imm_left if imm_left is not None else imm_right
if isinstance(imm, int) and 0 < imm <= 65535:
# 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"]
# shapes vary: dict or string
if isinstance(v, dict):
t = v.get("type") or v.get("kind")
if t:
action["type"] = t
# redirect / queue extra fields vary
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
# other expression types intentionally ignored for stateless reconstruction
if "type" not in action:
# 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 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}")
except RuntimeError:
logger.debug("no table %s.%s found when listing rules", family, table)
return []
# attempt to parse JSON output (binding can return textual JSON)
try:
data = json.loads(out)
except Exception:
# 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 listing failed: %s", err)
return []
data = json.loads(out_json)
results: List[Dict[str, Any]] = []
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("expressions", []) or []
recon = _reconstruct_rule_from_exprs(exprs)
# 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": 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
# ---------------------- API endpoints ----------------------
@router.get("/rules")
def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", family, table, chain)
try:
ensure_table_chain(family, table, chain)
except RuntimeError as e:
logger.error("failed to ensure table/chain %s.%s/%s: %s", family, table, chain, e)
raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {e}")
rules = list_rules_from_nft(family, table)
return {"count": len(rules), "rules": rules, "version": _current_version}
@router.put("/rules", response_model=ReplaceResult)
def put_rules(
rules: List[RuleModel],
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 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 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}")
if r.table and r.table != table:
raise HTTPException(status_code=400, detail=f"rule table mismatch: {r.table} != {table}")
if r.chain and r.chain != chain:
raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}")
# ensure table/chain exist first
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 script lines (flush + ordered add rules)
script_lines: List[str] = []
script_lines.append(f"flush chain {family} {table} {chain}")
for r in rules:
# assign client-side id if missing (not stored)
if not r.id:
r.id = str(uuid.uuid4())
# build match fragments
match_frag: List[str] = []
if r.match.iif:
match_frag += ["iif", f'"{r.match.iif}"']
if r.match.oif:
match_frag += ["oif", f'"{r.match.oif}"']
if r.match.meta_length is not None:
match_frag += ["meta", "length", str(r.match.meta_length)]
if r.match.ip_proto:
match_frag += ["ip", "protocol", str(r.match.ip_proto)]
if r.match.tcp_dport:
match_frag += ["tcp", "dport", str(r.match.tcp_dport)]
if r.match.udp_dport:
match_frag += ["udp", "dport", str(r.match.udp_dport)]
# build action fragment
action_frag: List[str] = []
a_type = r.action.type or "accept"
if a_type == "drop":
action_frag = ["drop"]
elif a_type == "accept":
action_frag = ["accept"]
elif a_type == "queue":
num = r.action.queue_num if r.action.queue_num is not None else 0
action_frag = ["queue", "num", str(num)]
elif a_type == "redirect":
if r.action.redirect_port is None:
raise HTTPException(status_code=400, detail="redirect action requires redirect_port")
action_frag = ["redirect", "to", f":{r.action.redirect_port}"]
else:
raise HTTPException(status_code=400, detail=f"unsupported action type: {a_type}")
parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag
script_lines.append(" ".join(parts))
# apply script lines via binding (line-by-line)
tmpfile_path: Optional[str] = None
try:
# keep a temp file copy for debugging if desired (optional)
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf:
tmpfile_path = tf.name
tf.write("\n".join(script_lines) + "\n")
tf.flush()
os.fsync(tf.fileno())
logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path)
try:
for ln in script_lines:
ln = ln.strip()
if not ln:
continue
nft_run_or_raise(ln)
except RuntimeError as e:
logger.error("failed applying nft script line '%s': %s", ln if 'ln' in locals() else "<unknown>", e)
raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}")
# success -> bump version
_current_version = str(uuid.uuid4())
logger.info("applied nft ruleset successfully; version=%s", _current_version)
return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))
finally:
if tmpfile_path and os.path.exists(tmpfile_path):
try:
os.remove(tmpfile_path)
except Exception:
logger.debug("failed to remove temp nft script %s", tmpfile_path)