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