All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
649 lines
24 KiB
Python
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))
|