test6
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s

This commit is contained in:
2026-01-11 13:47:26 +01:00
parent fa33f65dcb
commit d3723352a7

View File

@@ -13,7 +13,7 @@ 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, Header, Request
from fastapi import APIRouter, HTTPException, Request
from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any, Union
import logging
@@ -23,7 +23,11 @@ from enum import Enum
import asyncio
import subprocess
import re
import os
from backend.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
@@ -34,18 +38,17 @@ except Exception:
NFTablesBinding = None # will fall back to subprocess wrapper
# Router + logger
router = APIRouter()
router = APIRouter(prefix="/nft", tags=["nftables"])
logger = logging.getLogger("nftables")
logger.debug("nftables router module loaded")
# Defaults and version token
# Defaults
DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge"
_current_version: Optional[str] = None
# ---------------------- Enums ----------------------
# ---------------------- Enums used locally ----------------------
class ActionType(str, Enum):
DROP = "drop"
ACCEPT = "accept"
@@ -53,12 +56,6 @@ class ActionType(str, Enum):
REDIRECT = "redirect"
class Protocol(str, Enum):
ICMP = "icmp"
TCP = "tcp"
UDP = "udp"
class Family(str, Enum):
BRIDGE = "bridge"
INET = "inet"
@@ -72,7 +69,7 @@ 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, Protocol, str]] = None
ip_proto: Optional[Union[int, str]] = None # accept number or name (string)
tcp_dport: Optional[int] = None
udp_dport: Optional[int] = None
@@ -94,7 +91,6 @@ class RuleModel(BaseModel):
class ReplaceResult(BaseModel):
version: str
applied: bool
rules_count: int
@@ -143,11 +139,7 @@ class NFTSubprocessWrapper:
class NFTBindingWrapper:
"""
Wrapper using pyroute2 NFTables binding, instantiated synchronously if possible.
If the binding requires async setup and the current process already has a running event loop,
this wrapper raises so the factory can fall back to subprocess wrapper.
"""
"""Wrapper using pyroute2 NFTables binding, instantiated synchronously if possible."""
def __init__(self, binding_cls):
self._binding = None
@@ -182,7 +174,6 @@ class NFTBindingWrapper:
if not self._constructed:
raise RuntimeError(f"failed to instantiate NFTables binding: {last_exc}")
# If the binding exposes an async setup coroutine, ensure we can run it synchronously.
setup_coro = getattr(self._binding, "setup_endpoint", None)
if setup_coro and asyncio.iscoroutinefunction(setup_coro):
if asyncio.get_event_loop().is_running():
@@ -217,10 +208,6 @@ class NFTBindingWrapper:
def make_nft_wrapper():
"""
Choose the best NFT wrapper: try binding if available and can be used synchronously;
otherwise fall back to subprocess wrapper.
"""
if NFTablesBinding is not None:
try:
if asyncio.get_event_loop().is_running():
@@ -235,7 +222,6 @@ def make_nft_wrapper():
return NFTSubprocessWrapper()
# instantiate wrapper at import time
NFTC = make_nft_wrapper()
logger.info("selected nft wrapper: %s", type(NFTC).__name__)
@@ -250,12 +236,16 @@ def build_match_frag(match: MatchModel) -> List[str]:
if match.meta_length is not None:
frag += ["meta", "length", str(match.meta_length)]
if match.ip_proto is not None:
if isinstance(match.ip_proto, Protocol):
frag += ["ip", "protocol", match.ip_proto.value]
elif isinstance(match.ip_proto, int):
# 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:
frag += ["ip", "protocol", str(match.ip_proto)]
# 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:
@@ -301,7 +291,7 @@ def parse_comment_from_exprs(exprs: List[Dict[str, Any]]) -> Optional[str]:
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: handles common expression shapes; if not present returns defaults.
Best-effort parsing. Protocol numbers are converted to IPProtocolEnum member names when possible.
"""
exprs = []
if "rule" in entry:
@@ -340,6 +330,7 @@ def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
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
@@ -353,24 +344,26 @@ def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
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 low == "icmp":
match["ip_proto"] = 1
elif low == "tcp":
match["ip_proto"] = 6
elif low == "udp":
match["ip_proto"] = 17
# 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:
try:
match["ip_proto"] = int(imm)
except Exception:
match["ip_proto"] = imm
if isinstance(imm, int):
if 0 < imm < 256:
match["ip_proto"] = imm
elif 0 < imm <= 65535:
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")
@@ -419,15 +412,6 @@ def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
# ---------------------- helpers for normalized output ----------------------
def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]:
"""
Convert various wrapper outputs into a list of nft 'entries'.
Handles shapes:
- {'nftables': [ ... ]}
- {'out': '<json text>'}
- list([...])
- dict (single entry)
- {'out': '<plain text>'} -> returned as [{'text': <...>}]
"""
if not out:
return []
if isinstance(out, dict):
@@ -443,9 +427,7 @@ def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]:
return parsed
return [parsed]
except Exception:
# textual fallback
return [{"text": text}]
# single dict -> wrap
return [out]
if isinstance(out, list):
return out
@@ -455,10 +437,6 @@ def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]:
def _parse_textual_chain_output(text: str, family: str, table: str, chain: str) -> List[Dict[str, Any]]:
"""
Parse textual chain dump for simple rules as a last-resort fallback.
Looks for lines inside the chain and extracts simple patterns.
"""
results: List[Dict[str, Any]] = []
lines = text.splitlines()
for raw in lines:
@@ -484,16 +462,13 @@ def _parse_textual_chain_output(text: str, family: str, table: str, chain: str)
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 == "icmp":
match["ip_proto"] = 1
elif proto == "tcp":
match["ip_proto"] = 6
elif proto == "udp":
match["ip_proto"] = 17
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:
try:
match["ip_proto"] = int(proto)
except Exception:
match["ip_proto"] = proto
m2 = re.search(r'\btcp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE)
@@ -553,17 +528,8 @@ def ensure_table_chain(family: Union[str, Family], table: str, chain: str):
def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEFAULT_CHAIN) -> List[Dict[str, Any]]:
"""
Return reconstructed rules from kernel state.
Strategy:
1. Try `list chain <family> <table> <chain> -a` (best for rules + handles).
2. If that fails or returns textual output, fall back to `list ruleset` and extract all 'rule' entries.
3. If structured JSON is unavailable, parse textual output as last resort.
"""
fam = family.value if isinstance(family, Family) else family
# primary: list chain -a
try:
raw = NFTC.list_chain(fam, table, chain)
logger.debug("raw list_chain output: %s", str(raw)[:2000])
@@ -582,23 +548,15 @@ def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEF
for item in entries:
if not item:
continue
# textual full-dump fallback
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
# structured rule entry
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
# some outputs embed table/chain with nested lists (from list ruleset)
# find nested 'rule' if present
if isinstance(item, dict):
# item may be { 'table': {...} } or { 'chain': {...} } etc.
# attempt to find inner rule keys recursively
if "nftables" in item and isinstance(item["nftables"], list):
for sub in item["nftables"]:
if isinstance(sub, dict) and "rule" in sub:
@@ -608,7 +566,6 @@ def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEF
reconstructed = reconstruct_rule_from_rule_entry(item)
results.append(reconstructed)
else:
# nothing rule-like found here; skip
pass
return results
@@ -616,7 +573,6 @@ def list_rules_from_nft(family: Union[str, Family], table: str, chain: str = DEF
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
# flush chain
try:
NFTC.run(f"flush chain {fam} {table} {chain}")
except Exception as e:
@@ -647,25 +603,19 @@ def get_rules(family: Optional[Union[str, Family]] = DEFAULT_FAMILY,
raise HTTPException(status_code=500, detail=str(e))
rules = list_rules_from_nft(fam, table, chain)
return {"count": len(rules), "rules": rules, "version": _current_version}
return {"count": len(rules), "rules": rules}
@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[Union[str, Family]] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE,
chain: Optional[str] = DEFAULT_CHAIN,
):
global _current_version
fam = family.value if isinstance(family, Family) else family
# optimistic concurrency
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
for r in rules:
r_family_val = r.family.value if isinstance(r.family, Family) else r.family
@@ -695,6 +645,5 @@ def put_rules(
logger.error("failed to apply rules: %s", e)
raise HTTPException(status_code=500, detail=str(e))
_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))
logger.info("applied nft ruleset successfully; rules_count=%d", len(rules))
return ReplaceResult(applied=True, rules_count=len(rules))