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. - 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. - 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 pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any, Union from typing import Optional, List, Dict, Any, Union
import logging import logging
@@ -23,7 +23,11 @@ from enum import Enum
import asyncio import asyncio
import subprocess import subprocess
import re 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 to import NFTables binding (various pyroute2 layouts)
try: try:
from pyroute2.nftables.main import NFTables as NFTablesBinding # type: ignore from pyroute2.nftables.main import NFTables as NFTablesBinding # type: ignore
@@ -34,18 +38,17 @@ except Exception:
NFTablesBinding = None # will fall back to subprocess wrapper NFTablesBinding = None # will fall back to subprocess wrapper
# Router + logger # Router + logger
router = APIRouter() router = APIRouter(prefix="/nft", tags=["nftables"])
logger = logging.getLogger("nftables") logger = logging.getLogger("nftables")
logger.debug("nftables router module loaded") logger.debug("nftables router module loaded")
# Defaults and version token # Defaults
DEFAULT_TABLE = "mitm_tbl" DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward" DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge" DEFAULT_FAMILY = "bridge"
_current_version: Optional[str] = None
# ---------------------- Enums ---------------------- # ---------------------- Enums used locally ----------------------
class ActionType(str, Enum): class ActionType(str, Enum):
DROP = "drop" DROP = "drop"
ACCEPT = "accept" ACCEPT = "accept"
@@ -53,12 +56,6 @@ class ActionType(str, Enum):
REDIRECT = "redirect" REDIRECT = "redirect"
class Protocol(str, Enum):
ICMP = "icmp"
TCP = "tcp"
UDP = "udp"
class Family(str, Enum): class Family(str, Enum):
BRIDGE = "bridge" BRIDGE = "bridge"
INET = "inet" INET = "inet"
@@ -72,7 +69,7 @@ class MatchModel(BaseModel):
iif: Optional[str] = None iif: Optional[str] = None
oif: Optional[str] = None oif: Optional[str] = None
meta_length: Optional[Any] = None # int or range string 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 tcp_dport: Optional[int] = None
udp_dport: Optional[int] = None udp_dport: Optional[int] = None
@@ -94,7 +91,6 @@ class RuleModel(BaseModel):
class ReplaceResult(BaseModel): class ReplaceResult(BaseModel):
version: str
applied: bool applied: bool
rules_count: int rules_count: int
@@ -143,11 +139,7 @@ class NFTSubprocessWrapper:
class NFTBindingWrapper: class NFTBindingWrapper:
""" """Wrapper using pyroute2 NFTables binding, instantiated synchronously if possible."""
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.
"""
def __init__(self, binding_cls): def __init__(self, binding_cls):
self._binding = None self._binding = None
@@ -182,7 +174,6 @@ class NFTBindingWrapper:
if not self._constructed: if not self._constructed:
raise RuntimeError(f"failed to instantiate NFTables binding: {last_exc}") 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) setup_coro = getattr(self._binding, "setup_endpoint", None)
if setup_coro and asyncio.iscoroutinefunction(setup_coro): if setup_coro and asyncio.iscoroutinefunction(setup_coro):
if asyncio.get_event_loop().is_running(): if asyncio.get_event_loop().is_running():
@@ -217,10 +208,6 @@ class NFTBindingWrapper:
def make_nft_wrapper(): 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: if NFTablesBinding is not None:
try: try:
if asyncio.get_event_loop().is_running(): if asyncio.get_event_loop().is_running():
@@ -235,7 +222,6 @@ def make_nft_wrapper():
return NFTSubprocessWrapper() return NFTSubprocessWrapper()
# instantiate wrapper at import time
NFTC = make_nft_wrapper() NFTC = make_nft_wrapper()
logger.info("selected nft wrapper: %s", type(NFTC).__name__) 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: if match.meta_length is not None:
frag += ["meta", "length", str(match.meta_length)] frag += ["meta", "length", str(match.meta_length)]
if match.ip_proto is not None: if match.ip_proto is not None:
if isinstance(match.ip_proto, Protocol): # accept numeric or protocol-name (string). For names try to accept upper or lower.
frag += ["ip", "protocol", match.ip_proto.value] if isinstance(match.ip_proto, int):
elif isinstance(match.ip_proto, int):
frag += ["ip", "protocol", str(match.ip_proto)] frag += ["ip", "protocol", str(match.ip_proto)]
else: 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: if match.tcp_dport:
frag += ["tcp", "dport", str(match.tcp_dport)] frag += ["tcp", "dport", str(match.tcp_dport)]
if match.udp_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]: 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. 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 = [] exprs = []
if "rule" in entry: 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 {} cmp_obj = e.get("cmp") or e.get("match") or {}
left = cmp_obj.get("left") left = cmp_obj.get("left")
right = cmp_obj.get("right") right = cmp_obj.get("right")
def _extract_immediate(x): def _extract_immediate(x):
if not x or not isinstance(x, dict): if not x or not isinstance(x, dict):
return None return None
@@ -353,24 +344,26 @@ def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
return val return val
return val return val
return None return None
imm = _extract_immediate(left) or _extract_immediate(right) 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): if isinstance(imm, str):
low = imm.lower() low = imm.lower()
if low == "icmp": # If it's a known enum member name, return the enum member name (uppercase)
match["ip_proto"] = 1 if low.upper() in IPProtocolEnum.__members__:
elif low == "tcp": match["ip_proto"] = low.upper()
match["ip_proto"] = 6 else:
elif low == "udp": # numeric-string?
match["ip_proto"] = 17 if imm.isdigit():
match["ip_proto"] = protocol_from_number(int(imm))
else: else:
try:
match["ip_proto"] = int(imm)
except Exception:
match["ip_proto"] = imm match["ip_proto"] = imm
if isinstance(imm, int): if isinstance(imm, int):
if 0 < imm < 256: if 0 <= imm <= 255:
match["ip_proto"] = imm match["ip_proto"] = protocol_from_number(imm)
elif 0 < imm <= 65535: else:
# might be a port; treat as tcp_dport if sensible
if 0 < imm <= 65535:
match.setdefault("tcp_dport", imm) match.setdefault("tcp_dport", imm)
elif "verdict" in e or "immediate" in e or "return" in e: elif "verdict" in e or "immediate" in e or "return" in e:
v = e.get("verdict") or e.get("return") or e.get("immediate") 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 ---------------------- # ---------------------- helpers for normalized output ----------------------
def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]: 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: if not out:
return [] return []
if isinstance(out, dict): if isinstance(out, dict):
@@ -443,9 +427,7 @@ def _normalize_nft_output(out: Any) -> List[Dict[str, Any]]:
return parsed return parsed
return [parsed] return [parsed]
except Exception: except Exception:
# textual fallback
return [{"text": text}] return [{"text": text}]
# single dict -> wrap
return [out] return [out]
if isinstance(out, list): if isinstance(out, list):
return out 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]]: 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]] = [] results: List[Dict[str, Any]] = []
lines = text.splitlines() lines = text.splitlines()
for raw in lines: 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) m = re.search(r'\bip\s+protocol\s+([A-Za-z0-9_+-]+)\b', line_no_comment, flags=re.IGNORECASE)
if m: if m:
proto = m.group(1).lower() proto = m.group(1).lower()
if proto == "icmp": if proto.isdigit():
match["ip_proto"] = 1 match["ip_proto"] = protocol_from_number(int(proto))
elif proto == "tcp": else:
match["ip_proto"] = 6 # if proto corresponds to an enum member, return its name uppercase, otherwise the raw string
elif proto == "udp": if proto.upper() in IPProtocolEnum.__members__:
match["ip_proto"] = 17 match["ip_proto"] = proto.upper()
else: else:
try:
match["ip_proto"] = int(proto)
except Exception:
match["ip_proto"] = proto match["ip_proto"] = proto
m2 = re.search(r'\btcp\s+dport\s+(\d+)\b', line_no_comment, flags=re.IGNORECASE) 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]]: 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 fam = family.value if isinstance(family, Family) else family
# primary: list chain -a
try: try:
raw = NFTC.list_chain(fam, table, chain) raw = NFTC.list_chain(fam, table, chain)
logger.debug("raw list_chain output: %s", str(raw)[:2000]) 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: for item in entries:
if not item: if not item:
continue continue
# textual full-dump fallback
if isinstance(item, dict) and "text" in item and isinstance(item["text"], str): if isinstance(item, dict) and "text" in item and isinstance(item["text"], str):
results.extend(_parse_textual_chain_output(item["text"], fam, table, chain)) results.extend(_parse_textual_chain_output(item["text"], fam, table, chain))
continue continue
# structured rule entry
if "rule" in item: if "rule" in item:
rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")} rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")}
reconstructed = reconstruct_rule_from_rule_entry(rule_entry) reconstructed = reconstruct_rule_from_rule_entry(rule_entry)
results.append(reconstructed) results.append(reconstructed)
continue continue
# some outputs embed table/chain with nested lists (from list ruleset)
# find nested 'rule' if present
if isinstance(item, dict): 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): if "nftables" in item and isinstance(item["nftables"], list):
for sub in item["nftables"]: for sub in item["nftables"]:
if isinstance(sub, dict) and "rule" in sub: 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) reconstructed = reconstruct_rule_from_rule_entry(item)
results.append(reconstructed) results.append(reconstructed)
else: else:
# nothing rule-like found here; skip
pass pass
return results 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): 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 fam = family.value if isinstance(family, Family) else family
# flush chain
try: try:
NFTC.run(f"flush chain {fam} {table} {chain}") NFTC.run(f"flush chain {fam} {table} {chain}")
except Exception as e: 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)) raise HTTPException(status_code=500, detail=str(e))
rules = list_rules_from_nft(fam, table, chain) 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) @router.put("/rules", response_model=ReplaceResult)
def put_rules( def put_rules(
rules: List[RuleModel], rules: List[RuleModel],
request: Request, request: Request,
if_match: Optional[str] = Header(None, alias="If-Match"),
family: Optional[Union[str, Family]] = DEFAULT_FAMILY, family: Optional[Union[str, Family]] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE, table: Optional[str] = DEFAULT_TABLE,
chain: Optional[str] = DEFAULT_CHAIN, chain: Optional[str] = DEFAULT_CHAIN,
): ):
global _current_version
fam = family.value if isinstance(family, Family) else family 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 # validate per-rule family/table/chain
for r in rules: for r in rules:
r_family_val = r.family.value if isinstance(r.family, Family) else r.family 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) logger.error("failed to apply rules: %s", e)
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))
_current_version = str(uuid.uuid4()) logger.info("applied nft ruleset successfully; rules_count=%d", len(rules))
logger.info("applied nft ruleset successfully; version=%s", _current_version) return ReplaceResult(applied=True, rules_count=len(rules))
return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))