This commit is contained in:
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user