improve models and use better nft lib
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
This commit is contained in:
@@ -1,28 +1,28 @@
|
||||
# fastapi_nft_stateless_comment_enums.py
|
||||
"""
|
||||
Stateless FastAPI nftables router that uses libnftables binding.
|
||||
|
||||
- GET /nft/rules -> reconstruct rule objects from nft kernel state (best-effort)
|
||||
- PUT /nft/rules -> replace entire ordered ruleset (applies via binding, line-by-line)
|
||||
Stateless FastAPI nftables router using pyroute2.nftables.
|
||||
|
||||
- Stateless: no in-process or on-disk rule store.
|
||||
- Rules may include an optional 'comment' field that will be written into nft's comment.
|
||||
- Enums introduced for Action.type, ip_proto (common names) and family.
|
||||
- Endpoints:
|
||||
- GET /nft/rules -> reconstruct rules from kernel (returns comment if present)
|
||||
- PUT /nft/rules -> replace entire ordered rule set (clients supply optional comment per rule)
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, Header, Request
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import Optional, List, Dict, Any
|
||||
import uuid
|
||||
import json
|
||||
import tempfile
|
||||
import os
|
||||
from typing import Optional, List, Dict, Any, Union
|
||||
import logging
|
||||
import json
|
||||
import uuid
|
||||
from enum import Enum
|
||||
|
||||
# try import libnftables binding
|
||||
try:
|
||||
from nftables import Nftables
|
||||
except Exception:
|
||||
Nftables = None # clearer error raised when attempting to use binding
|
||||
# pyroute2 nftables
|
||||
from pyroute2 import nftables
|
||||
|
||||
router = APIRouter(prefix="/nft", tags=["nftables"])
|
||||
logger = logging.getLogger("nftables")
|
||||
logger.debug("nftables stateless router module loaded")
|
||||
logger.debug("nftables stateless-comment-enums router loaded")
|
||||
|
||||
# Defaults and version token
|
||||
DEFAULT_TABLE = "mitm_tbl"
|
||||
@@ -30,30 +30,54 @@ DEFAULT_CHAIN = "forward"
|
||||
DEFAULT_FAMILY = "bridge"
|
||||
_current_version: Optional[str] = None
|
||||
|
||||
|
||||
# ---------------------- Enums ----------------------
|
||||
class ActionType(str, Enum):
|
||||
DROP = "drop"
|
||||
ACCEPT = "accept"
|
||||
QUEUE = "queue"
|
||||
REDIRECT = "redirect"
|
||||
|
||||
|
||||
class Protocol(str, Enum):
|
||||
ICMP = "icmp"
|
||||
TCP = "tcp"
|
||||
UDP = "udp"
|
||||
|
||||
|
||||
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 like "100-200"
|
||||
ip_proto: Optional[Any] = None # numeric or name
|
||||
meta_length: Optional[Any] = None # int or range string
|
||||
# ip_proto may be Protocol enum, integer or arbitrary string (name)
|
||||
ip_proto: Optional[Union[int, Protocol, str]] = None
|
||||
tcp_dport: Optional[int] = None
|
||||
udp_dport: Optional[int] = None
|
||||
|
||||
|
||||
class ActionModel(BaseModel):
|
||||
type: Optional[str] = None # drop | accept | queue | redirect
|
||||
type: ActionType # ActionType enum
|
||||
queue_num: Optional[int] = None
|
||||
redirect_port: Optional[int] = None
|
||||
|
||||
|
||||
class RuleModel(BaseModel):
|
||||
# id is optional client-side convenience only (not persisted)
|
||||
id: Optional[str] = Field(None, description="optional client id; not stored server-side")
|
||||
family: Optional[str] = Field(DEFAULT_FAMILY)
|
||||
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):
|
||||
@@ -62,216 +86,210 @@ class ReplaceResult(BaseModel):
|
||||
rules_count: int
|
||||
|
||||
|
||||
# ---------------------- nft binding helpers ----------------------
|
||||
def _ensure_binding_available():
|
||||
if Nftables is None:
|
||||
raise RuntimeError(
|
||||
"python nftables binding not available. Install system package `python3-nftables` or a compatible binding."
|
||||
)
|
||||
# ---------------------- nft wrapper ----------------------
|
||||
class NFT:
|
||||
def __init__(self):
|
||||
self.nft = nftables.NFTables()
|
||||
|
||||
|
||||
def nft_cmd(cmd: str):
|
||||
"""
|
||||
Run a libnftables command string.
|
||||
Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error.
|
||||
"""
|
||||
_ensure_binding_available()
|
||||
nft = Nftables()
|
||||
try:
|
||||
rc, out, err = nft.cmd(cmd)
|
||||
# ensure we work with strings
|
||||
if isinstance(out, bytes):
|
||||
out = out.decode(errors="ignore")
|
||||
if isinstance(err, bytes):
|
||||
err = err.decode(errors="ignore")
|
||||
return rc, out or "", err or ""
|
||||
except Exception as e:
|
||||
logger.exception("unexpected nft binding exception for cmd=%s", cmd)
|
||||
raise RuntimeError(f"nft binding error: {e}")
|
||||
|
||||
|
||||
def nft_run_or_raise(cmd: str) -> str:
|
||||
"""
|
||||
Run command via binding and raise RuntimeError if rc != 0.
|
||||
Returns stdout string on success.
|
||||
"""
|
||||
rc, out, err = nft_cmd(cmd)
|
||||
logger.debug("nft cmd: %s -> rc=%s out_len=%d err_len=%d", cmd, rc, len(out), len(err))
|
||||
if rc != 0:
|
||||
# surface err if available, otherwise a generic message
|
||||
raise RuntimeError(err or f"nft command '{cmd}' failed (rc={rc})")
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------- chain/table ensure / reconstruction ----------------------
|
||||
def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
||||
"""
|
||||
Ensure the nft table and chain exist. Use libnftables `add` commands first; on failure
|
||||
attempt to apply a tiny script by invoking individual commands (no "-f" with the binding).
|
||||
Raises RuntimeError on persistent failure.
|
||||
"""
|
||||
logger.info("ensuring table %s.%s exists", family, table)
|
||||
|
||||
# try add table
|
||||
try:
|
||||
nft_run_or_raise(f"add table {family} {table}")
|
||||
logger.debug("created table %s.%s via add table", family, table)
|
||||
except RuntimeError as e_table:
|
||||
logger.info("add table failed: %s. attempting fallback", e_table)
|
||||
# fallback: build script and apply the lines individually
|
||||
script_lines = [f"table {family} {table} {{ }}"]
|
||||
def run(self, cmd: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Run a single nft command string via pyroute2.NFTables.cmd.
|
||||
Returns parsed JSON if possible, otherwise a dict with 'out' textual output.
|
||||
Raises RuntimeError on failure.
|
||||
"""
|
||||
try:
|
||||
for ln in script_lines:
|
||||
nft_run_or_raise(ln)
|
||||
logger.debug("created table %s.%s via fallback lines", family, table)
|
||||
except RuntimeError as e2:
|
||||
logger.error("fallback for creating table failed: %s (original: %s)", e2, e_table)
|
||||
raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2
|
||||
|
||||
# try add chain
|
||||
logger.info("ensuring chain %s in table %s exists", chain, table)
|
||||
try:
|
||||
nft_run_or_raise(
|
||||
f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}'
|
||||
)
|
||||
logger.debug("created chain %s in %s.%s via add chain", chain, family, table)
|
||||
except RuntimeError as e_chain:
|
||||
logger.info("add chain failed: %s. attempting fallback", e_chain)
|
||||
script_lines = [
|
||||
f"table {family} {table} {{",
|
||||
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}",
|
||||
f"}}",
|
||||
]
|
||||
try:
|
||||
for ln in script_lines:
|
||||
nft_run_or_raise(ln)
|
||||
logger.debug("created chain %s in %s.%s via fallback lines", chain, family, table)
|
||||
except RuntimeError as e2:
|
||||
logger.error("fallback for creating chain failed: %s (original: %s)", e2, e_chain)
|
||||
raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2
|
||||
|
||||
logger.info("table/chain ensured: %s.%s/%s", family, table, chain)
|
||||
|
||||
|
||||
# ---------------------- reconstruct rules from nft JSON (best-effort) ----------------------
|
||||
# Replace your previous _reconstruct_rule_from_exprs and list_rules_from_nft with these.
|
||||
|
||||
def _extract_immediate_value(node):
|
||||
"""Return an integer/string value from various immediate/data shapes, or None."""
|
||||
if not node or not isinstance(node, dict):
|
||||
return None
|
||||
# common keys used by libnftables JSON
|
||||
for k in ("immediate", "value", "data", "s", "v"):
|
||||
if k in node:
|
||||
val = node[k]
|
||||
# hex strings like "0x00000050"
|
||||
if isinstance(val, str) and val.startswith("0x"):
|
||||
rc, out, err = self.nft.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 int(val, 16)
|
||||
return json.loads(out)
|
||||
except Exception:
|
||||
return val
|
||||
return val
|
||||
# sometimes immediate sits under {"immediate": {"value": ...}}
|
||||
if "right" in node or "left" in node:
|
||||
# caller handles left/right structures
|
||||
return None
|
||||
return {"out": out}
|
||||
return {}
|
||||
except AttributeError:
|
||||
# Try json_cmd fallback if older pyroute2 version
|
||||
out = self.nft.json_cmd(cmd)
|
||||
return out or {}
|
||||
except Exception as e:
|
||||
logger.exception("nft wrapper error for cmd=%s: %s", cmd, e)
|
||||
raise
|
||||
|
||||
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}")
|
||||
|
||||
|
||||
NFTC = NFT()
|
||||
|
||||
|
||||
# ---------------------- 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:
|
||||
# ip_proto can be Protocol enum, int or str
|
||||
if isinstance(match.ip_proto, Protocol):
|
||||
frag += ["ip", "protocol", match.ip_proto.value]
|
||||
elif isinstance(match.ip_proto, int):
|
||||
frag += ["ip", "protocol", str(match.ip_proto)]
|
||||
else:
|
||||
frag += ["ip", "protocol", str(match.ip_proto)]
|
||||
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]:
|
||||
# ActionModel.type is an ActionType enum
|
||||
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:
|
||||
"""
|
||||
Build fragment after 'add rule <family> <table> <chain>'.
|
||||
Includes comment if rule.comment provided.
|
||||
"""
|
||||
match_frag = build_match_frag(rule.match)
|
||||
action_frag = build_action_frag(rule.action)
|
||||
parts = match_frag + action_frag
|
||||
if rule.comment:
|
||||
# include comment as-is (user-provided). Wrap in quotes.
|
||||
parts += ['comment', f'"{rule.comment}"']
|
||||
return " ".join(parts)
|
||||
|
||||
|
||||
def parse_comment_from_exprs(exprs: List[Dict[str, Any]]) -> Optional[str]:
|
||||
"""Scan expressions for a comment expression and return its string if present."""
|
||||
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_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
def reconstruct_rule_from_rule_entry(entry: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Improved best-effort mapping from nft expression JSON to our RuleModel fields.
|
||||
This function attempts to match many of the shapes produced by libnftables.
|
||||
Given a rule entry returned by NFTC.list_table, return a dict that resembles RuleModel
|
||||
(family/table/chain + match + action + comment if present). This is best-effort parsing.
|
||||
"""
|
||||
# normalize expressions
|
||||
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)
|
||||
|
||||
# best-effort reconstruction of match & action
|
||||
match: Dict[str, Any] = {}
|
||||
action: Dict[str, Any] = {}
|
||||
|
||||
# Track last payload hint (if present) so that subsequent cmp can be interpreted
|
||||
last_payload_hint: Optional[Dict[str, Any]] = None
|
||||
|
||||
for idx, e in enumerate(exprs):
|
||||
# handle meta (iif/oif/length)
|
||||
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")
|
||||
# sometimes v is dict
|
||||
if isinstance(v, dict):
|
||||
v = _extract_immediate_value(v) or v.get("s") or v.get("v")
|
||||
if key in ("iifname", "iif", "in", "iifname?"):
|
||||
if v:
|
||||
match["iif"] = v
|
||||
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"):
|
||||
if v:
|
||||
match["oif"] = v
|
||||
match["oif"] = v
|
||||
elif key in ("length", "len"):
|
||||
if v is not None:
|
||||
match["meta_length"] = v
|
||||
else:
|
||||
logger.debug("meta with unknown key: %s value=%s", key, v)
|
||||
|
||||
# store payload hints for later comparisons
|
||||
elif "payload" in e:
|
||||
last_payload_hint = e["payload"]
|
||||
# make it easier for cmp handling: include index
|
||||
last_payload_hint["_idx"] = idx
|
||||
|
||||
# cmp (compare) expressions: left/right might be payload / immediate structures
|
||||
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")
|
||||
# extract immediate numeric if present
|
||||
imm_left = _extract_immediate_value(left)
|
||||
imm_right = _extract_immediate_value(right)
|
||||
imm = imm_left if imm_left is not None else imm_right
|
||||
|
||||
# If either immediate is a small int treat as ip_proto
|
||||
if isinstance(imm, int) and 0 < imm < 256:
|
||||
# prefer to store as number (frontend may show numeric)
|
||||
match["ip_proto"] = imm
|
||||
|
||||
# If immediate looks like port (1-65535), try to detect target (tcp/udp)
|
||||
if isinstance(imm, int) and 0 < imm <= 65535:
|
||||
assigned = False
|
||||
# heuristics: if left/right contains a payload hint referring to TCP/UDP or 'dport' strings:
|
||||
for side in (left, right):
|
||||
if isinstance(side, dict):
|
||||
# payload form used by libnftables can include 'protocol' or 'field'
|
||||
pl = side.get("payload") or side.get("left", {}).get("payload")
|
||||
if isinstance(pl, dict):
|
||||
protocol_hint = pl.get("protocol") or pl.get("proto") or pl.get("family")
|
||||
field_hint = pl.get("field") or pl.get("meta")
|
||||
sh = json.dumps(pl).lower()
|
||||
if "tcp" in sh or "sport" in sh or "dport" in sh or "th" in sh:
|
||||
match["tcp_dport"] = imm
|
||||
assigned = True
|
||||
break
|
||||
if "udp" in sh or "udph" in sh or "udp." in sh:
|
||||
match["udp_dport"] = imm
|
||||
assigned = True
|
||||
break
|
||||
if not assigned:
|
||||
# fallback: if ip_proto already indicates tcp(6) or udp(17), assign accordingly
|
||||
proto = match.get("ip_proto")
|
||||
if proto in (6, "tcp"):
|
||||
match["tcp_dport"] = imm
|
||||
elif proto in (17, "udp"):
|
||||
match["udp_dport"] = imm
|
||||
else:
|
||||
# if we can't know, default to tcp_dport as most common case
|
||||
match.setdefault("tcp_dport", imm)
|
||||
|
||||
# verdict / immediate verdict expressions -> action
|
||||
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 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
|
||||
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:
|
||||
match.setdefault("tcp_dport", imm)
|
||||
elif "verdict" in e or "immediate" in e or "return" in e:
|
||||
# verdict may be string or dict
|
||||
v = e.get("verdict") or e.get("return") or e.get("immediate")
|
||||
# normalize dict forms
|
||||
if isinstance(v, dict):
|
||||
t = v.get("type") or v.get("kind") or v.get("verdict")
|
||||
if t:
|
||||
action["type"] = t
|
||||
# redirect/queue shaped differently across versions
|
||||
if "to" in v:
|
||||
action["type"] = "redirect"
|
||||
action["redirect_port"] = v.get("to")
|
||||
@@ -279,103 +297,128 @@ def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
action["type"] = "queue"
|
||||
action["queue_num"] = v.get("queue")
|
||||
elif isinstance(v, str):
|
||||
# strings like "accept" or "drop"
|
||||
action["type"] = v
|
||||
else:
|
||||
# sometimes verdict is expressed as nested dict under 'verdict': {'kind':'accept'}
|
||||
if isinstance(e.get("verdict"), dict):
|
||||
vv = e["verdict"]
|
||||
action["type"] = vv.get("kind") or vv.get("type")
|
||||
|
||||
# old-style 'match' entries with left/right payload/immediate
|
||||
elif "match" in e:
|
||||
m = e["match"]
|
||||
left = m.get("left")
|
||||
right = m.get("right")
|
||||
imm_left = _extract_immediate_value(left)
|
||||
imm_right = _extract_immediate_value(right)
|
||||
if isinstance(imm_left, int) and 0 < imm_left < 256:
|
||||
match["ip_proto"] = imm_left
|
||||
if isinstance(imm_right, int) and 0 < imm_right < 256:
|
||||
match["ip_proto"] = imm_right
|
||||
|
||||
else:
|
||||
# unknown expression type: log (DEBUG) for later tuning
|
||||
logger.debug("unhandled nft expr type: %s", list(e.keys()))
|
||||
# ignore other expression types
|
||||
pass
|
||||
|
||||
# default action to accept if kernel default or none found
|
||||
if "type" not in action:
|
||||
action["type"] = "accept"
|
||||
|
||||
return {"match": match, "action": action}
|
||||
|
||||
|
||||
def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Reconstruct rules using libnftables JSON output with more robust parsing.
|
||||
"""
|
||||
# ensure binding present and request JSON output explicitly
|
||||
_ensure_binding_available()
|
||||
nft = Nftables()
|
||||
nft.set_json_output(True)
|
||||
rc, out, err = nft.cmd(f"list table {family} {table}")
|
||||
if rc != 0:
|
||||
logger.debug("nft list table returned rc=%s err=%s", rc, err)
|
||||
return []
|
||||
|
||||
# convert action.type to ActionType enum if possible
|
||||
try:
|
||||
data = json.loads(out)
|
||||
action_type_val = action.get("type")
|
||||
if isinstance(action_type_val, str):
|
||||
action["type"] = ActionType(action_type_val)
|
||||
except Exception:
|
||||
# leave as-is if conversion fails
|
||||
pass
|
||||
|
||||
# convert 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,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------- 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.error("failed to parse nft JSON output: %s", 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) -> List[Dict[str, Any]]:
|
||||
fam = family.value if isinstance(family, Family) else family
|
||||
try:
|
||||
out = NFTC.list_table(fam, table)
|
||||
except Exception as e:
|
||||
logger.debug("list_table failed: %s", e)
|
||||
return []
|
||||
|
||||
entries = []
|
||||
if isinstance(out, dict) and "nftables" in out:
|
||||
entries = out["nftables"]
|
||||
elif isinstance(out, dict) and "out" in out and isinstance(out["out"], str):
|
||||
try:
|
||||
parsed = json.loads(out["out"])
|
||||
if isinstance(parsed, dict) and "nftables" in parsed:
|
||||
entries = parsed["nftables"]
|
||||
elif isinstance(parsed, list):
|
||||
entries = parsed
|
||||
except Exception:
|
||||
logger.debug("could not parse textual nft output")
|
||||
entries = []
|
||||
elif isinstance(out, list):
|
||||
entries = out
|
||||
elif isinstance(out, dict) and out:
|
||||
entries = [out]
|
||||
else:
|
||||
entries = []
|
||||
|
||||
results: List[Dict[str, Any]] = []
|
||||
for item in data.get("nftables", []):
|
||||
if "rule" not in item:
|
||||
for item in entries:
|
||||
if not item:
|
||||
continue
|
||||
rule_obj = item["rule"]
|
||||
# different versions might use 'expr', 'expressions' or 'expr'
|
||||
exprs = rule_obj.get("expr") or rule_obj.get("exprs") or rule_obj.get("expressions") or []
|
||||
# if exprs is not a list but a dict (some shaped outputs) normalize
|
||||
if isinstance(exprs, dict):
|
||||
# sometimes expressions are nested inside a single dict; try to find array keys
|
||||
for k in ("expr", "expressions", "exprs"):
|
||||
v = exprs.get(k)
|
||||
if isinstance(v, list):
|
||||
exprs = v
|
||||
break
|
||||
else:
|
||||
# give up and wrap
|
||||
exprs = [exprs]
|
||||
if "rule" in item:
|
||||
rule_entry = {"rule": item["rule"], "family": item.get("family"), "table": item.get("table")}
|
||||
else:
|
||||
rule_entry = item
|
||||
reconstructed = reconstruct_rule_from_rule_entry(rule_entry)
|
||||
results.append(reconstructed)
|
||||
|
||||
recon = _reconstruct_rule_from_exprs(exprs)
|
||||
chain_name = rule_obj.get("chain") or rule_obj.get("chain_name") or rule_obj.get("table") or DEFAULT_CHAIN
|
||||
# make output match your RuleModel-ish structure (no id)
|
||||
rule_dict = {
|
||||
"family": family,
|
||||
"table": table,
|
||||
"chain": chain_name,
|
||||
"match": recon.get("match", {}),
|
||||
"action": recon.get("action", {}),
|
||||
}
|
||||
results.append(rule_dict)
|
||||
|
||||
logger.debug("reconstructed %d rules from %s.%s", len(results), family, table)
|
||||
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
|
||||
# flush chain
|
||||
try:
|
||||
NFTC.run(f"flush chain {fam} {table} {chain}")
|
||||
except Exception as e:
|
||||
logger.debug("flush chain may have returned error: %s", e)
|
||||
|
||||
# add rules
|
||||
for r in rules:
|
||||
# ensure family value converted to string
|
||||
fam_r = r.family.value if isinstance(r.family, Family) else r.family
|
||||
frag = nft_rule_fragment_from_model(r)
|
||||
try:
|
||||
NFTC.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[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
|
||||
logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", family, table, chain)
|
||||
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(family, table, chain)
|
||||
except RuntimeError as e:
|
||||
logger.error("failed to ensure table/chain %s.%s/%s: %s", family, table, chain, e)
|
||||
raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {e}")
|
||||
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(family, table)
|
||||
rules = list_rules_from_nft(fam, table)
|
||||
return {"count": len(rules), "rules": rules, "version": _current_version}
|
||||
|
||||
|
||||
@@ -384,107 +427,46 @@ def put_rules(
|
||||
rules: List[RuleModel],
|
||||
request: Request,
|
||||
if_match: Optional[str] = Header(None, alias="If-Match"),
|
||||
family: Optional[str] = DEFAULT_FAMILY,
|
||||
family: Optional[Union[str, Family]] = DEFAULT_FAMILY,
|
||||
table: Optional[str] = DEFAULT_TABLE,
|
||||
chain: Optional[str] = DEFAULT_CHAIN,
|
||||
):
|
||||
"""
|
||||
Replace entire ordered ruleset. Validates provided rules, then builds nft commands
|
||||
and applies them line-by-line via libnftables binding (no '-f' token passed into binding).
|
||||
"""
|
||||
# optimistic concurrency
|
||||
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 (quick checks)
|
||||
# validate per-rule family/table/chain
|
||||
for r in rules:
|
||||
if r.family and r.family != family:
|
||||
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}")
|
||||
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 first
|
||||
# ensure table/chain exist
|
||||
try:
|
||||
ensure_table_chain(family, table, chain)
|
||||
except RuntimeError as e:
|
||||
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))
|
||||
|
||||
# build script lines (flush + ordered add rules)
|
||||
script_lines: List[str] = []
|
||||
script_lines.append(f"flush chain {family} {table} {chain}")
|
||||
|
||||
# ensure rule ids for client convenience
|
||||
for r in rules:
|
||||
# assign client-side id if missing (not stored)
|
||||
if not r.id:
|
||||
r.id = str(uuid.uuid4())
|
||||
|
||||
# build match fragments
|
||||
match_frag: List[str] = []
|
||||
if r.match.iif:
|
||||
match_frag += ["iif", f'"{r.match.iif}"']
|
||||
if r.match.oif:
|
||||
match_frag += ["oif", f'"{r.match.oif}"']
|
||||
if r.match.meta_length is not None:
|
||||
match_frag += ["meta", "length", str(r.match.meta_length)]
|
||||
if r.match.ip_proto:
|
||||
match_frag += ["ip", "protocol", str(r.match.ip_proto)]
|
||||
if r.match.tcp_dport:
|
||||
match_frag += ["tcp", "dport", str(r.match.tcp_dport)]
|
||||
if r.match.udp_dport:
|
||||
match_frag += ["udp", "dport", str(r.match.udp_dport)]
|
||||
|
||||
# build action fragment
|
||||
action_frag: List[str] = []
|
||||
a_type = r.action.type or "accept"
|
||||
if a_type == "drop":
|
||||
action_frag = ["drop"]
|
||||
elif a_type == "accept":
|
||||
action_frag = ["accept"]
|
||||
elif a_type == "queue":
|
||||
num = r.action.queue_num if r.action.queue_num is not None else 0
|
||||
action_frag = ["queue", "num", str(num)]
|
||||
elif a_type == "redirect":
|
||||
if r.action.redirect_port is None:
|
||||
raise HTTPException(status_code=400, detail="redirect action requires redirect_port")
|
||||
action_frag = ["redirect", "to", f":{r.action.redirect_port}"]
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail=f"unsupported action type: {a_type}")
|
||||
|
||||
parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag
|
||||
script_lines.append(" ".join(parts))
|
||||
|
||||
# apply script lines via binding (line-by-line)
|
||||
tmpfile_path: Optional[str] = None
|
||||
# attempt to replace rules
|
||||
try:
|
||||
# keep a temp file copy for debugging if desired (optional)
|
||||
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf:
|
||||
tmpfile_path = tf.name
|
||||
tf.write("\n".join(script_lines) + "\n")
|
||||
tf.flush()
|
||||
os.fsync(tf.fileno())
|
||||
logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path)
|
||||
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))
|
||||
|
||||
try:
|
||||
for ln in script_lines:
|
||||
ln = ln.strip()
|
||||
if not ln:
|
||||
continue
|
||||
nft_run_or_raise(ln)
|
||||
except RuntimeError as e:
|
||||
logger.error("failed applying nft script line '%s': %s", ln if 'ln' in locals() else "<unknown>", e)
|
||||
raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}")
|
||||
|
||||
# success -> bump version
|
||||
_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))
|
||||
finally:
|
||||
if tmpfile_path and os.path.exists(tmpfile_path):
|
||||
try:
|
||||
os.remove(tmpfile_path)
|
||||
except Exception:
|
||||
logger.debug("failed to remove temp nft script %s", tmpfile_path)
|
||||
_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))
|
||||
|
||||
Reference in New Issue
Block a user