Files
mitm-webserver/backend/src/api/nftables_api.py
malmert 439712479c
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
fixifix
2026-01-11 13:48:23 +01:00

649 lines
24 KiB
Python

# fastapi_nft_stateless_comment_enums_resilient_full.py
"""
Stateless FastAPI nftables router with enums and resilient pyroute2 binding.
- Stateless: no in-process or on-disk rule store.
- Rules may include an optional 'comment' field that will be written into nft's comment.
- Uses pyroute2 binding if available and can be instantiated synchronously, otherwise falls back to calling the `nft` CLI via subprocess.
- Exposes two endpoints under /nft:
GET /nft/rules -> list rules reconstructed from kernel state (returns comment if present)
PUT /nft/rules -> replace entire ordered ruleset (clients supply optional comment per rule)
Notes:
- The service must run with privileges to modify nftables (root / CAP_NET_ADMIN) when applying rules.
- If the subprocess fallback is used, ensure the `nft` binary is present.
"""
from fastapi import APIRouter, HTTPException, Request
from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any, Union
import logging
import json
import uuid
from enum import Enum
import asyncio
import subprocess
import re
from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number
# ---------------------- main module (resilient wrapper + API) ----------------------
# Try to import NFTables binding (various pyroute2 layouts)
try:
from pyroute2.nftables.main import NFTables as NFTablesBinding # type: ignore
except Exception:
try:
from pyroute2.nftables import NFTables as NFTablesBinding # type: ignore
except Exception:
NFTablesBinding = None # will fall back to subprocess wrapper
# Router + logger
router = APIRouter()
logger = logging.getLogger("nftables")
logger.debug("nftables router module loaded")
# Defaults
DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge"
# ---------------------- Enums used locally ----------------------
class ActionType(str, Enum):
DROP = "drop"
ACCEPT = "accept"
QUEUE = "queue"
REDIRECT = "redirect"
class Family(str, Enum):
BRIDGE = "bridge"
INET = "inet"
IP = "ip"
IP6 = "ip6"
ARP = "arp"
# ---------------------- Pydantic models ----------------------
class MatchModel(BaseModel):
iif: Optional[str] = None
oif: Optional[str] = None
meta_length: Optional[Any] = None # int or range string
ip_proto: Optional[Union[int, str]] = None # accept number or name (string)
tcp_dport: Optional[int] = None
udp_dport: Optional[int] = None
class ActionModel(BaseModel):
type: ActionType
queue_num: Optional[int] = None
redirect_port: Optional[int] = None
class RuleModel(BaseModel):
id: Optional[str] = Field(None, description="optional client id; not persisted")
family: Optional[Family] = Field(Family.BRIDGE)
table: Optional[str] = Field(DEFAULT_TABLE)
chain: Optional[str] = Field(DEFAULT_CHAIN)
match: MatchModel
action: ActionModel
comment: Optional[str] = Field(None, description="optional human-readable comment stored in nft comment")
class ReplaceResult(BaseModel):
applied: bool
rules_count: int
# ---------------------- resilient NFT wrapper selection ----------------------
class NFTSubprocessWrapper:
"""Fallback wrapper calling the `nft` CLI via subprocess."""
def __init__(self):
try:
subprocess.run(["nft", "--version"], check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
logger.debug("nft CLI available for subprocess wrapper")
except Exception as e:
logger.warning("nft CLI not available or couldn't be invoked: %s", e)
def run(self, cmd: str) -> Dict[str, Any]:
args = ["nft"] + cmd.split()
try:
proc = subprocess.run(args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True)
out = proc.stdout or ""
try:
return json.loads(out)
except Exception:
return {"out": out}
except subprocess.CalledProcessError as e:
logger.error("nft CLI failed cmd=%s stderr=%s", cmd, e.stderr)
raise RuntimeError(e.stderr or str(e))
def add_table(self, family: str, table: str):
return self.run(f"add table {family} {table}")
def add_chain(self, family: str, table: str, chain: str, type_: str = "filter", hook: str = "forward", priority: int = 0, policy: str = "accept"):
return self.run(f'add chain {family} {table} {chain} {{ type {type_} hook {hook} priority {priority}; policy {policy}; }}')
def list_table(self, family: str, table: str):
return self.run(f"list table {family} {table}")
def list_chain(self, family: str, table: str, chain: str):
return self.run(f"list chain {family} {table} {chain} -a")
def add_rule(self, family: str, table: str, chain: str, rule_fragment: str):
return self.run(f"add rule {family} {table} {chain} {rule_fragment}")
def delete_rule_by_handle(self, family: str, table: str, chain: str, handle: str):
return self.run(f"delete rule {family} {table} {chain} handle {handle}")
class NFTBindingWrapper:
"""Wrapper using pyroute2 NFTables binding, instantiated synchronously if possible."""
def __init__(self, binding_cls):
self._binding = None
self._constructed = False
tried_kwargs = [
{"async": False},
{"nl_async": False},
{"use_async": False},
{"asyncio": False},
]
last_exc = None
for kw in tried_kwargs:
try:
self._binding = binding_cls(**kw)
self._constructed = True
logger.debug("NFTables binding instantiated with kwargs %s", kw)
break
except TypeError as e:
last_exc = e
except Exception as e:
last_exc = e
if not self._constructed:
try:
self._binding = binding_cls()
self._constructed = True
logger.debug("NFTables binding instantiated with no kwargs")
except Exception as e:
last_exc = e
if not self._constructed:
raise RuntimeError(f"failed to instantiate NFTables binding: {last_exc}")
setup_coro = getattr(self._binding, "setup_endpoint", None)
if setup_coro and asyncio.iscoroutinefunction(setup_coro):
if asyncio.get_event_loop().is_running():
raise RuntimeError("pyroute2 NFTables requires async setup but event loop is already running")
try:
asyncio.get_event_loop().run_until_complete(setup_coro())
except Exception as e:
logger.error("awaiting binding.setup_endpoint() failed: %s", e)
raise
def run(self, cmd: str) -> Dict[str, Any]:
try:
rc, out, err = self._binding.cmd(cmd)
if isinstance(out, bytes):
out = out.decode(errors="ignore")
if isinstance(err, bytes):
err = err.decode(errors="ignore")
if rc != 0:
raise RuntimeError(err or f"nft cmd failed rc={rc}")
if out:
try:
return json.loads(out)
except Exception:
return {"out": out}
return {}
except AttributeError:
out = self._binding.json_cmd(cmd)
return out or {}
except Exception as e:
logger.exception("nft binding cmd failed: %s", e)
raise
def make_nft_wrapper():
if NFTablesBinding is not None:
try:
if asyncio.get_event_loop().is_running():
logger.info("asyncio loop is running; skipping binding and using subprocess wrapper")
raise RuntimeError("event loop running")
w = NFTBindingWrapper(NFTablesBinding)
logger.info("using pyroute2 NFTables binding")
return w
except Exception as e:
logger.warning("pyroute2 binding unavailable/synchronous construction failed: %s; falling back to nft CLI", e)
logger.info("using nft CLI subprocess wrapper")
return NFTSubprocessWrapper()
NFTC = make_nft_wrapper()
logger.info("selected nft wrapper: %s", type(NFTC).__name__)
# ---------------------- builders / parsers ----------------------
def build_match_frag(match: MatchModel) -> List[str]:
frag: List[str] = []
if match.iif:
frag += ["iif", f'"{match.iif}"']
if match.oif:
frag += ["oif", f'"{match.oif}"']
if match.meta_length is not None:
frag += ["meta", "length", str(match.meta_length)]
if match.ip_proto is not None:
# accept numeric or protocol-name (string). For names try to accept upper or lower.
if isinstance(match.ip_proto, int):
frag += ["ip", "protocol", str(match.ip_proto)]
else:
# if it's a known IPProtocolEnum name, use lowercase for nft syntax
v = str(match.ip_proto)
if v.upper() in IPProtocolEnum.__members__:
frag += ["ip", "protocol", v.lower()]
else:
frag += ["ip", "protocol", v]
if match.tcp_dport:
frag += ["tcp", "dport", str(match.tcp_dport)]
if match.udp_dport:
frag += ["udp", "dport", str(match.udp_dport)]
return frag
def build_action_frag(action: ActionModel) -> List[str]:
if action.type == ActionType.DROP:
return ["drop"]
if action.type == ActionType.ACCEPT:
return ["accept"]
if action.type == ActionType.QUEUE:
num = action.queue_num if action.queue_num is not None else 0
return ["queue", "num", str(num)]
if action.type == ActionType.REDIRECT:
if action.redirect_port is None:
raise ValueError("redirect action requires redirect_port")
return ["redirect", "to", f":{action.redirect_port}"]
raise ValueError(f"unsupported action type: {action.type}")
def nft_rule_fragment_from_model(rule: RuleModel) -> str:
match_frag = build_match_frag(rule.match)
action_frag = build_action_frag(rule.action)
parts = match_frag + action_frag
if rule.comment:
parts += ['comment', f'"{rule.comment}"']
return " ".join(parts)
def parse_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
def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
"""
Convert pyroute2/nft `rule` JSON entry into the simple dict format returned by the API.
Best-effort parsing. Protocol numbers are converted to IPProtocolEnum member names when possible.
"""
exprs = []
if "rule" in entry:
r = entry["rule"]
exprs = r.get("expr") or r.get("expressions") or r.get("exprs") or []
chain_name = r.get("chain") or entry.get("chain") or r.get("chain_name")
family = entry.get("family") or DEFAULT_FAMILY
table = entry.get("table") or DEFAULT_TABLE
else:
exprs = entry.get("expr") or entry.get("expressions") or entry.get("exprs") or []
chain_name = entry.get("chain") or entry.get("chain_name")
family = entry.get("family") or DEFAULT_FAMILY
table = entry.get("table") or DEFAULT_TABLE
if isinstance(exprs, dict):
exprs = [exprs]
comment = parse_comment_from_exprs(exprs)
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") or m.get("field")
v = m.get("v") or m.get("s") or m.get("value")
if isinstance(v, dict):
v = v.get("value") or v.get("v") or v.get("s")
if key in ("iifname", "iif", "in"):
match["iif"] = v
elif key in ("oifname", "oif", "out"):
match["oif"] = v
elif key in ("length", "len"):
match["meta_length"] = v
elif "cmp" in e or "match" in e:
cmp_obj = e.get("cmp") or e.get("match") or {}
left = cmp_obj.get("left")
right = cmp_obj.get("right")
def _extract_immediate(x):
if not x or not isinstance(x, dict):
return None
for k in ("immediate", "value", "data", "s", "v"):
if k in x:
val = x[k]
if isinstance(val, str) and val.startswith("0x"):
try:
return int(val, 16)
except Exception:
return val
return val
return None
imm = _extract_immediate(left) or _extract_immediate(right)
# If imm is a string name (like 'icmp'), keep it; if int then map to IPProtocolEnum name if possible.
if isinstance(imm, str):
low = imm.lower()
# If it's a known enum member name, return the enum member name (uppercase)
if low.upper() in IPProtocolEnum.__members__:
match["ip_proto"] = low.upper()
else:
# numeric-string?
if imm.isdigit():
match["ip_proto"] = protocol_from_number(int(imm))
else:
match["ip_proto"] = imm
if isinstance(imm, int):
if 0 <= imm <= 255:
match["ip_proto"] = protocol_from_number(imm)
else:
# might be a port; treat as tcp_dport if sensible
if 0 < imm <= 65535:
match.setdefault("tcp_dport", imm)
elif "verdict" in e or "immediate" in e or "return" in e:
v = e.get("verdict") or e.get("return") or e.get("immediate")
if isinstance(v, dict):
t = v.get("type") or v.get("kind") or v.get("verdict")
if t:
action["type"] = t
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
else:
# ignore other expression types
pass
if "type" not in action:
action["type"] = "accept"
# try to coerce action.type to ActionType
try:
if isinstance(action.get("type"), str):
action["type"] = ActionType(action["type"])
except Exception:
pass
# coerce family to Family enum if possible
try:
if isinstance(family, str):
family = Family(family)
except Exception:
pass
return {
"family": family,
"table": table,
"chain": chain_name or DEFAULT_CHAIN,
"match": match,
"action": action,
"comment": comment,
}
# ---------------------- helpers for normalized output ----------------------
def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]:
if not out:
return []
if isinstance(out, dict):
if "nftables" in out and isinstance(out["nftables"], list):
return out["nftables"]
if "out" in out and isinstance(out["out"], str):
text = out["out"]
try:
parsed = json.loads(text)
if isinstance(parsed, dict) and "nftables" in parsed:
return parsed["nftables"]
if isinstance(parsed, list):
return parsed
return [parsed]
except Exception:
return [{"text": text}]
return [out]
if isinstance(out, list):
return out
if isinstance(out, str):
return [{"text": out}]
return []
def _parse_textual_chain_output(text: str, family: str, table: str, chain: str) -> List[Dict[str, Any]]:
results: List[Dict[str, Any]] = []
lines = text.splitlines()
for raw in lines:
line = raw.strip()
if not line:
continue
if line.startswith("type ") or line.startswith("policy ") or line.startswith("chain ") or line.startswith("table "):
continue
if line in ("{", "}"):
continue
comment = None
cm = re.search(r'comment\s+"([^"]+)"', line)
if cm:
comment = cm.group(1)
line_no_comment = re.sub(r'comment\s+"[^"]+"', '', line)
else:
line_no_comment = line
match: Dict[str, Any] = {}
action: Dict[str, Any] = {"type": "accept"}
m = re.search(r'\bip\s+protocol\s+([A-Za-z0-9_+-]+)\b', line_no_comment, flags=re.IGNORECASE)
if m:
proto = m.group(1).lower()
if proto.isdigit():
match["ip_proto"] = protocol_from_number(int(proto))
else:
# if proto corresponds to an enum member, return its name uppercase, otherwise the raw string
if proto.upper() in IPProtocolEnum.__members__:
match["ip_proto"] = proto.upper()
else:
match["ip_proto"] = proto
m2 = re.search(r'\btcp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE)
if m2:
match["tcp_dport"] = int(m2.group(1))
m3 = re.search(r'\budp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE)
if m3:
match["udp_dport"] = int(m3.group(1))
m4 = re.search(r'\biif\s+"?([^"\s]+)"?\b', line_no_comment, flags=re.IGNORECASE)
if m4:
match["iif"] = m4.group(1)
m5 = re.search(r'\boif\s+"?([^"\s]+)"?\b', line_no_comment, flags=re.IGNORECASE)
if m5:
match["oif"] = m5.group(1)
if re.search(r'\bdrop\b', line_no_comment, flags=re.IGNORECASE):
action = {"type": "drop"}
elif re.search(r'\baccept\b', line_no_comment, flags=re.IGNORECASE):
action = {"type": "accept"}
elif re.search(r'\bqueue\b', line_no_comment, flags=re.IGNORECASE):
q = re.search(r'queue\s+num\s+(\d+)', line_no_comment, flags=re.IGNORECASE)
if q:
action = {"type": "queue", "queue_num": int(q.group(1))}
else:
action = {"type": "queue"}
elif re.search(r'\bredirect\b', line_no_comment, flags=re.IGNORECASE) or re.search(r'\bto\s+:\d+\b', line_no_comment, flags=re.IGNORECASE):
mredir = re.search(r':(\d+)', line_no_comment)
if mredir:
action = {"type": "redirect", "redirect_port": int(mredir.group(1))}
else:
action = {"type": "redirect"}
results.append({
"family": family,
"table": table,
"chain": chain,
"match": match,
"action": action,
"comment": comment,
})
return results
# ---------------------- high-level operations ----------------------
def ensure_table_chain(family: Union[str, Family], table: str, chain: str):
fam = family.value if isinstance(family, Family) else family
try:
NFTC.add_table(fam, table)
except Exception as e:
logger.debug("add_table may have failed/exists: %s", e)
try:
NFTC.add_chain(fam, table, chain)
except Exception as e:
logger.debug("add_chain may have failed/exists: %s", e)
def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEFAULT_CHAIN) -> List[Dict[str, Any]]:
fam = family.value if isinstance(family, Family) else family
try:
raw = NFTC.list_chain(fam, table, chain)
logger.debug("raw list_chain output: %s", str(raw)[:2000])
except Exception as e:
logger.debug("list_chain failed (%s); falling back to list ruleset", e)
try:
raw = NFTC.run("list ruleset")
logger.debug("raw list ruleset output: %s", str(raw)[:2000])
except Exception as e2:
logger.error("list ruleset failed: %s", e2)
return []
entries = _normalize_nft_output(raw)
results: List[Dict[str, Any]] = []
for item in entries:
if not item:
continue
if isinstance(item, dict) and "text" in item and isinstance(item["text"], str):
results.extend(_parse_textual_chain_output(item["text"], fam, table, chain))
continue
if "rule" in item:
rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")}
reconstructed = reconstruct_rule_from_rule_entry(rule_entry)
results.append(reconstructed)
continue
if isinstance(item, dict):
if "nftables" in item and isinstance(item["nftables"], list):
for sub in item["nftables"]:
if isinstance(sub, dict) and "rule" in sub:
reconstructed = reconstruct_rule_from_rule_entry(sub)
results.append(reconstructed)
elif item.get("type") == "rule" or "expr" in item or "expressions" in item:
reconstructed = reconstruct_rule_from_rule_entry(item)
results.append(reconstructed)
else:
pass
return results
def add_rules_replace_all(rules: List[RuleModel], family: Union[str, Family], table: str, chain: str):
fam = family.value if isinstance(family, Family) else family
try:
NFTC.run(f"flush chain {fam} {table} {chain}")
except Exception as e:
logger.debug("flush chain may have returned error: %s", e)
for r in rules:
fam_r = r.family.value if isinstance(r.family, Family) else r.family
frag = nft_rule_fragment_from_model(r)
try:
NFTC.run(f"add rule {fam_r} {r.table} {r.chain} {frag}")
logger.info("added rule frag=%s", frag)
except Exception as e:
logger.exception("failed to add rule: %s", e)
raise RuntimeError(f"failed to add rule: {e}")
# ---------------------- API endpoints ----------------------
@router.get("/rules")
def get_rules(family: Optional[Union[str, Family]] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE,
chain: Optional[str] = DEFAULT_CHAIN):
fam = family.value if isinstance(family, Family) else family
logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", fam, table, chain)
try:
ensure_table_chain(fam, table, chain)
except Exception as e:
logger.error("failed to ensure table/chain: %s", e)
raise HTTPException(status_code=500, detail=str(e))
rules = list_rules_from_nft(fam, table, chain)
return {"count": len(rules), "rules": rules}
@router.put("/rules", response_model=ReplaceResult)
def put_rules(
rules: List[RuleModel],
request: Request,
family: Optional[Union[str, Family]] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE,
chain: Optional[str] = DEFAULT_CHAIN,
):
fam = family.value if isinstance(family, Family) else family
# validate per-rule family/table/chain
for r in rules:
r_family_val = r.family.value if isinstance(r.family, Family) else r.family
if r.family and r_family_val != fam:
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r_family_val} != {fam}")
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
try:
ensure_table_chain(fam, table, chain)
except Exception as e:
logger.error("failed to ensure table/chain: %s", e)
raise HTTPException(status_code=500, detail=str(e))
# ensure rule ids for convenience
for r in rules:
if not r.id:
r.id = str(uuid.uuid4())
# attempt to replace rules
try:
add_rules_replace_all(rules, fam, table, chain)
except Exception as e:
logger.error("failed to apply rules: %s", e)
raise HTTPException(status_code=500, detail=str(e))
logger.info("applied nft ruleset successfully; rules_count=%d", len(rules))
return ReplaceResult(applied=True, rules_count=len(rules))