All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s
407 lines
16 KiB
Python
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)
|