Files
mitm-webserver/backend/src/api/nft_api.py
malmert 20c77739ba
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
fix trailing
2026-01-11 17:59:08 +01:00

538 lines
19 KiB
Python

# fastapi_nft_router.py
# -*- coding: utf-8 -*-
"""
FastAPI router that lists and manages nftables rules for family 'bridge'
(default table 'mitm_tbl', chain 'forward').
Behavior:
- Uses `nft --json list ruleset` to obtain authoritative rule metadata (handles).
- Uses `nft list chain <family> <table> <chain>` to extract the exact textual
rule lines. Mapping is done by matching `handle N` in the textual output.
- Returns for each rule:
- nft_rule_text_full: exact line from 'nft list chain ...' including 'handle N' (or None)
- nft_rule_text: same line trimmed to remove trailing 'handle N' (or None)
- add_command: "add rule <table> <chain> <nft_rule_text>" (or None)
- No JSON->text reconstruction is attempted. If text mapping is missing we return None.
Security note:
- Process must be run with privileges to run nft (root or appropriate capabilities).
- Consider adding auth before exposing these endpoints.
"""
from typing import Any, Dict, List, Optional, Union, Literal
import subprocess
import shutil
import json
import logging
import re
from enum import Enum
from fastapi import APIRouter, HTTPException, Body
from pydantic import BaseModel, Field, validator
# Router & logging
router = APIRouter()
logger = logging.getLogger("nftables")
logger.debug("nftables router module loaded")
# Defaults & nft binary
DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge"
NFT_BIN = shutil.which("nft")
# ----------------- Enums (for frontend) -----------------
class Family(str, Enum):
bridge = DEFAULT_FAMILY
class Table(str, Enum):
table = DEFAULT_TABLE
class Chain(str, Enum):
forward = "forward"
input = "input"
output = "output"
class MetaKey(str, Enum):
iifname = "iifname"
oifname = "oifname"
iif = "iif"
oif = "oif"
prio = "prio"
class EtherField(str, Enum):
saddr = "saddr"
daddr = "daddr"
class IPDir(str, Enum):
saddr = "saddr"
daddr = "daddr"
class Op(str, Enum):
eq = "=="
neq = "!="
lt = "<"
gt = ">"
contains = "in"
class Verdict(str, Enum):
accept = "accept"
drop = "drop"
reject = "reject"
continue_ = "continue"
class Proto(str, Enum):
tcp = "tcp"
udp = "udp"
icmp = "icmp"
class ConntrackState(str, Enum):
new = "new"
established = "established"
related = "related"
invalid = "invalid"
class RejectType(str, Enum):
icmp = "icmp"
tcp_reset = "tcp reset"
class IcmpType(str, Enum):
dest_unreachable = "destination-unreachable"
time_exceeded = "time-exceeded"
echo_reply = "echo-reply"
echo_request = "echo-request"
port_unreachable = "port-unreachable"
host_unreachable = "host-unreachable"
fragmentation_needed = "fragmentation-needed"
class LogGroup(int, Enum):
g0 = 0
g1 = 1
g2 = 2
g3 = 3
g4 = 4
g5 = 5
g6 = 6
g7 = 7
# ------------ Pydantic expression models (typed for frontend) ----------
class BaseExpr(BaseModel):
kind: str
class Config:
extra = "forbid"
class MetaExpr(BaseExpr):
kind: Literal["meta"] = Field(default="meta")
key: MetaKey
op: Op = Op.eq
value: str
class EtherExpr(BaseExpr):
kind: Literal["ether"] = Field(default="ether")
field: EtherField
op: Op = Op.eq
value: str
class IPExpr(BaseExpr):
kind: Literal["ip"] = Field(default="ip")
side: IPDir
op: Op = Op.eq
value: str
class ProtoPortExpr(BaseExpr):
kind: Literal["l4"] = Field(default="l4")
proto: Proto
sport: Optional[str] = None
dport: Optional[str] = None
class CTEexpr(BaseExpr):
kind: Literal["ct"] = Field(default="ct")
state: ConntrackState
class VerdictExpr(BaseExpr):
kind: Literal["verdict"] = Field(default="verdict")
verdict: Verdict
class RejectExpr(BaseExpr):
kind: Literal["reject"] = Field(default="reject")
reject_type: RejectType
icmp_type: Optional[IcmpType] = None
class LogExpr(BaseExpr):
kind: Literal["log"] = Field(default="log")
prefix: Optional[str] = None
group: Optional[LogGroup] = None
class RawExpr(BaseExpr):
kind: Literal["raw"] = Field(default="raw")
snippet: str
Expr = Union[
MetaExpr, EtherExpr, IPExpr, ProtoPortExpr, CTEexpr,
VerdictExpr, RejectExpr, LogExpr, RawExpr,
]
# ---------------- Rule model ----------------
class RuleModel(BaseModel):
family: Family = Family.bridge
table: Table = Table.table
chain: Chain = Chain.forward
expr: List[Expr] = Field(default_factory=list)
comment: Optional[str] = None
position: Optional[int] = None # 1-based
handle: Optional[int] = None
@validator("family")
def only_bridge(cls, v: Family) -> Family:
if v != Family.bridge:
raise ValueError("This router only manages family 'bridge'")
return v
# ----------------- Helpers --------------------
def ensure_nft_available() -> None:
if not NFT_BIN:
logger.error("nft binary not found on server")
raise HTTPException(status_code=500, detail="nft binary not found on server")
def run_nft_cmd(cmd: str) -> Dict[str, Any]:
"""
Execute a single nft script line via `nft -f -`. Returns stdout/stderr.
"""
ensure_nft_available()
full_cmd = [NFT_BIN, "-f", "-"]
script = cmd.rstrip() + "\n"
logger.info("Running nft command: %s", cmd)
logger.debug("Exec: %s ; script: %s", full_cmd, script)
try:
proc = subprocess.run(full_cmd, input=script.encode(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=True)
stdout = proc.stdout.decode()
stderr = proc.stderr.decode()
logger.info("nft success (stdout %d bytes, stderr %d bytes)", len(stdout), len(stderr))
logger.debug("nft stdout: %s", stdout or "<empty>")
if stderr:
logger.debug("nft stderr: %s", stderr)
return {"stdout": stdout, "stderr": stderr}
except subprocess.CalledProcessError as e:
err = e.stderr.decode() if e.stderr else str(e)
logger.error("nft failed: %s", err)
raise HTTPException(status_code=500, detail=err)
def ensure_table_and_chain_exist(family: str, table: str, chain: str) -> None:
"""
Ensure the named table and chain exist; create them with conservative defaults if missing.
"""
logger.debug("Ensure table/chain exist family=%s table=%s chain=%s", family, table, chain)
ensure_nft_available()
try:
out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE)
parsed = json.loads(out)
except subprocess.CalledProcessError as e:
logger.error("Failed to list ruleset: %s", e.stderr.decode())
raise HTTPException(status_code=500, detail=f"nft failed: {e.stderr.decode()}")
items = parsed.get("nftables") if isinstance(parsed, dict) else parsed
if not isinstance(items, list):
items = []
table_exists = False
chain_exists = False
for it in items:
if "table" in it:
t = it["table"]
if isinstance(t, dict) and t.get("name") == table and t.get("family") == family:
table_exists = True
if "chain" in it:
ch = it["chain"]
if isinstance(ch, dict) and ch.get("name") == chain and ch.get("table") == table and ch.get("family") == family:
chain_exists = True
if not table_exists:
logger.info("Creating table %s %s", family, table)
run_nft_cmd(f"add table {family} {table}")
if not chain_exists:
logger.info("Creating chain %s in table %s", chain, table)
if chain in ("input", "forward", "output"):
run_nft_cmd(f"add chain {family} {table} {chain} {{ type filter hook {chain} priority 0; policy accept; }}")
else:
run_nft_cmd(f"add chain {family} {table} {chain} {{ policy accept; }}")
# ----------------- Text mapping (strict) -----------------
HANDLE_RE = re.compile(r"\bhandle\s+(\d+)\b", flags=re.IGNORECASE)
def build_handle_text_map(family: str, table: str, chain: str) -> Dict[int, str]:
"""
Runs: nft list chain <family> <table> <chain>
Returns mapping handle -> full textual line containing 'handle N'.
If the textual output cannot be retrieved, raises HTTPException.
"""
ensure_nft_available()
cmd = [NFT_BIN, "--handle", "list", "chain", family, table, chain]
logger.debug("Listing chain text: %s", " ".join(cmd))
try:
out = subprocess.check_output(cmd, stderr=subprocess.PIPE).decode()
except subprocess.CalledProcessError as e:
logger.error("Failed to list chain text: %s", e.stderr.decode())
raise HTTPException(status_code=500, detail=e.stderr.decode())
mapping: Dict[int, str] = {}
for line in out.splitlines():
s = line.strip()
if not s:
continue
m = HANDLE_RE.search(s)
if not m:
continue
try:
h = int(m.group(1))
# full textual line as-is
mapping[h] = s
logger.debug("Found textual rule for handle %d: %s", h, s)
except Exception as ex:
logger.debug("Failed parsing handle from line: %s (%s)", s, ex)
continue
return mapping
# ----------------- Rules listing (JSON + strict textual lookup) -----------------
def nft_list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
"""
Return rules parsed from nft --json list ruleset, augmented with textual lines
extracted from `nft list chain <family> <table> <chain>` via handle matching.
For each rule returned:
- family, table, chain
- handle
- position (1-based in chain)
- comment (best-effort from JSON exprs)
- verdict (best-effort)
- exprs (the JSON expr list)
- nft_rule_text_full: exact textual line from nft list chain ... INCLUDING 'handle N' (or None)
- nft_rule_text: textual line trimmed to remove trailing 'handle N' (or None)
- add_command: "add rule <table> <chain> <nft_rule_text>" (or None)
"""
ensure_nft_available()
# 1) JSON dump: authoritative structure
try:
out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE)
parsed = json.loads(out)
except subprocess.CalledProcessError as e:
logger.error("Failed to get JSON ruleset: %s", e.stderr.decode())
raise HTTPException(status_code=500, detail=e.stderr.decode())
# 2) textual map: strict mapping by handle
text_map: Dict[int, str] = {}
try:
text_map = build_handle_text_map(DEFAULT_FAMILY, table, chain)
logger.debug("Text map size: %d", len(text_map))
except HTTPException as e:
# bubble up the error: user asked to extract exact textual lines and we couldn't get them
logger.error("Failed to obtain textual chain dump: %s", getattr(e, "detail", str(e)))
# still continue — per your request we won't attempt reconstructions, but we can return None textual fields.
text_map = {}
results: List[Dict[str, Any]] = []
counters: Dict[str, int] = {}
items = parsed.get("nftables") if isinstance(parsed, dict) else parsed
if not isinstance(items, list):
items = []
for it in items:
if "rule" not in it:
continue
r = it["rule"]
family = r.get("family")
table_name = r.get("table")
chain_name = r.get("chain")
# only return rules for requested table/chain
if table_name != table or chain_name != chain:
continue
key = f"{family}:{table_name}:{chain_name}"
counters.setdefault(key, 0)
counters[key] += 1
position = counters[key]
handle = r.get("handle")
exprs = r.get("expr", [])
# best-effort comment + verdict extraction from JSON exprs (keeps UI useful)
comment: Optional[str] = None
verdict: Optional[str] = None
for ex in exprs:
if not isinstance(ex, dict):
continue
if "comment" in ex:
c = ex.get("comment")
if isinstance(c, str):
comment = c
elif isinstance(c, dict):
comment = c.get("text") or c.get("str")
if "verdict" in ex:
v = ex["verdict"]
if isinstance(v, dict):
verdict = next(iter(v.keys()), None)
else:
verdict = str(v)
if "drop" in ex and verdict is None:
verdict = "drop"
if "accept" in ex and verdict is None:
verdict = "accept"
if "reject" in ex and verdict is None:
verdict = "reject"
# strict textual lookup: only use exact line if present in text_map
nft_rule_text_full: Optional[str] = None
nft_rule_text: Optional[str] = None
add_command: Optional[str] = None
if handle is not None and handle in text_map:
nft_rule_text_full = text_map[handle]
# remove trailing ' handle N' to get copy/paste clause
m = HANDLE_RE.search(nft_rule_text_full)
if m:
# slice everything before ' handle N'
raw_clause = nft_rule_text_full[: m.start()].strip()
else:
raw_clause = nft_rule_text_full
# remove a trailing lone '#' (and surrounding whitespace) if present
# e.g. "meta iifname \"eth0\" # " -> "meta iifname \"eth0\""
nft_rule_text = re.sub(r"\s*#\s*$", "", raw_clause).strip()
add_command = f"add rule {table_name} {chain_name} {nft_rule_text}".strip() if nft_rule_text else None
logger.debug("Attached textual rule for handle %s", handle)
else:
# If mapping missing, per your instruction do not attempt to reconstruct — leave textual fields None
logger.debug("No textual mapping for handle %s — textual fields will be None", handle)
results.append({
"family": family,
"table": table_name,
"chain": chain_name,
"handle": handle,
"position": position,
"comment": comment,
"verdict": verdict,
"exprs": exprs,
"nft_rule_text_full": nft_rule_text_full,
"nft_rule_text": nft_rule_text,
"add_command": add_command,
})
return {"rules": results}
# ------------ Expr -> nft snippet & command builder (preview/add) ------
def expr_to_nft_snippet(e: Expr) -> str:
"""Build short nft snippet from typed Expr (used for preview/add)."""
if isinstance(e, MetaExpr):
val = e.value
key = e.key.value
return f"meta {key} {e.op.value} {val}"
if isinstance(e, EtherExpr):
return f"ether {e.field.value} {e.op.value} {e.value}"
if isinstance(e, IPExpr):
return f"ip {e.side.value} {e.op.value} {e.value}"
if isinstance(e, ProtoPortExpr):
parts = [e.proto.value]
if e.sport:
parts.append(f"sport {e.sport}")
if e.dport:
parts.append(f"dport {e.dport}")
return " ".join(parts)
if isinstance(e, CTEexpr):
return f"ct state {e.state.value}"
if isinstance(e, VerdictExpr):
return e.verdict.value if e.verdict != Verdict.continue_ else "continue"
if isinstance(e, RejectExpr):
if e.reject_type == RejectType.icmp:
return f"reject with icmp type {e.icmp_type.value}" if e.icmp_type else "reject"
return e.reject_type.value
if isinstance(e, LogExpr):
parts = ["log"]
if e.prefix:
parts.append(f'prefix "{e.prefix}"')
if e.group is not None:
parts.append(f"group {int(e.group)}")
return " ".join(parts)
if isinstance(e, RawExpr):
return e.snippet
raise ValueError("Unsupported expression type")
def rule_to_nft_cmd(rule: RuleModel) -> str:
expr_snips = [expr_to_nft_snippet(e) for e in rule.expr]
body = " ".join(s for s in expr_snips if s)
if rule.position is not None:
cmd = f"insert rule {rule.table.value} {rule.chain.value} position {rule.position} {body}"
else:
cmd = f"add rule {rule.table.value} {rule.chain.value} {body}"
if rule.comment:
cmd += f' comment "{rule.comment}"'
return cmd
# ---------------- Endpoints -------------------
@router.get("/options")
def get_options() -> Dict[str, Any]:
"""Return enum choices for frontend dropdowns."""
return {
"family": [f.value for f in Family],
"table": [t.value for t in Table],
"chain": [c.value for c in Chain],
"meta_keys": [m.value for m in MetaKey],
"ether_fields": [e.value for e in EtherField],
"ip_dirs": [d.value for d in IPDir],
"ops": [o.value for o in Op],
"verdicts": [v.value for v in Verdict],
"protocols": [p.value for p in Proto],
"ct_states": [s.value for s in ConntrackState],
"reject_types": [r.value for r in RejectType],
"icmp_types": [i.value for i in IcmpType],
"log_groups": [int(g.value) for g in LogGroup],
}
@router.get("/rules")
def list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
"""List rules for the given table/chain (ensures table/chain exist first)."""
ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain)
return nft_list_rules(table=table, chain=chain)
@router.post("/rules/preview")
def preview_rule(rule: RuleModel = Body(...)) -> Dict[str, str]:
"""Return the nft command that would be executed for the provided rule (preview only)."""
try:
cmd = rule_to_nft_cmd(rule)
except Exception as e:
logger.error("Preview build failed: %s", e)
raise HTTPException(status_code=400, detail=str(e))
return {"cmd": cmd}
@router.post("/rules")
def add_rule(rule: RuleModel = Body(...)) -> Dict[str, Any]:
"""Insert/append rule (creates table/chain if missing)."""
ensure_table_and_chain_exist(rule.family.value, rule.table.value, rule.chain.value)
cmd = rule_to_nft_cmd(rule)
return run_nft_cmd(cmd)
@router.delete("/rules/{handle}")
def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
"""Delete rule by nft handle."""
ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain)
cmd = f"delete rule {table} {chain} handle {handle}"
return run_nft_cmd(cmd)
@router.put("/rules/{handle}")
def update_rule(handle: int, rule: RuleModel = Body(...), table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
"""Replace a rule by handle: delete by handle then insert replacement (attempt to preserve position)."""
ensure_table_and_chain_exist(rule.family.value, rule.table.value, rule.chain.value)
rules_info = nft_list_rules(table=table, chain=chain)
position: Optional[int] = None
for r in rules_info.get("rules", []):
if r.get("handle") == handle:
position = r.get("position")
break
delete_rule(handle, table=table, chain=chain)
if position is not None:
rule.position = position
return add_rule(rule)