From d3723352a77daa30696e109845772a889111ce46 Mon Sep 17 00:00:00 2001 From: malmert Date: Sun, 11 Jan 2026 13:47:26 +0100 Subject: [PATCH] test6 --- backend/src/api/nftables_api.py | 139 ++++++++++---------------------- 1 file changed, 44 insertions(+), 95 deletions(-) diff --git a/backend/src/api/nftables_api.py b/backend/src/api/nftables_api.py index c3e43fc..105bbbe 100644 --- a/backend/src/api/nftables_api.py +++ b/backend/src/api/nftables_api.py @@ -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,25 +344,27 @@ 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: - try: - match["ip_proto"] = int(imm) - except Exception: + # 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 < 256: - match["ip_proto"] = imm - elif 0 < imm <= 65535: - match.setdefault("tcp_dport", imm) + 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): @@ -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': ''} - - list([...]) - - dict (single entry) - - {'out': ''} -> 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: - try: - match["ip_proto"] = int(proto) - except Exception: + # 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) @@ -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 -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))