This commit is contained in:
@@ -1,46 +1,41 @@
|
|||||||
# fastapi_nft_stateless.py
|
|
||||||
"""
|
"""
|
||||||
Stateless FastAPI nftables router.
|
Stateless FastAPI nftables router that uses libnftables binding.
|
||||||
Only uses kernel-stored info (nft) as the source of truth.
|
|
||||||
|
|
||||||
Endpoints (mounted under /nft):
|
- GET /nft/rules -> reconstruct rule objects from nft kernel state (best-effort)
|
||||||
- GET /rules -> reconstruct rule objects from nft kernel state (best-effort)
|
- PUT /nft/rules -> replace entire ordered ruleset (applies via binding, line-by-line)
|
||||||
- PUT /rules -> replace entire ordered ruleset (applies via nft -f)
|
|
||||||
|
|
||||||
No persistence, no comments used for mapping.
|
|
||||||
"""
|
"""
|
||||||
from fastapi import APIRouter, HTTPException, Header, Request
|
from fastapi import APIRouter, HTTPException, Header, Request
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
from typing import Optional, List, Dict, Any, Union, Tuple
|
from typing import Optional, List, Dict, Any
|
||||||
import uuid
|
import uuid
|
||||||
import json
|
import json
|
||||||
import tempfile
|
import tempfile
|
||||||
import os
|
import os
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
# libnftables binding
|
# try import libnftables binding
|
||||||
try:
|
try:
|
||||||
from nftables import Nftables
|
from nftables import Nftables
|
||||||
except Exception:
|
except Exception:
|
||||||
Nftables = None # will raise when used
|
Nftables = None # clearer error raised when attempting to use binding
|
||||||
|
|
||||||
router = APIRouter(prefix="/nft", tags=["nftables"])
|
router = APIRouter(prefix="/nft", tags=["nftables"])
|
||||||
logger = logging.getLogger("nftables")
|
logger = logging.getLogger("nftables")
|
||||||
logger.debug("nftables stateless router loaded")
|
logger.debug("nftables stateless router module loaded")
|
||||||
|
|
||||||
# Defaults
|
# Defaults and version token
|
||||||
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
|
_current_version: Optional[str] = None
|
||||||
|
|
||||||
# ---------------------- Pydantic models ----------------------
|
# ---------------------- Pydantic models ----------------------
|
||||||
class MatchModel(BaseModel):
|
class MatchModel(BaseModel):
|
||||||
iif: Optional[str] = None
|
iif: Optional[str] = None
|
||||||
oif: Optional[str] = None
|
oif: Optional[str] = None
|
||||||
meta_length: Optional[Union[int, str]] = None
|
meta_length: Optional[Any] = None # int or range string like "100-200"
|
||||||
ip_proto: Optional[Union[int, str]] = None
|
ip_proto: Optional[Any] = None # numeric or name
|
||||||
tcp_dport: Optional[int] = None
|
tcp_dport: Optional[int] = None
|
||||||
udp_dport: Optional[int] = None
|
udp_dport: Optional[int] = None
|
||||||
|
|
||||||
@@ -52,8 +47,8 @@ class ActionModel(BaseModel):
|
|||||||
|
|
||||||
|
|
||||||
class RuleModel(BaseModel):
|
class RuleModel(BaseModel):
|
||||||
# note: since we don't persist IDs, id remains optional
|
# id is optional client-side convenience only (not persisted)
|
||||||
id: Optional[str] = Field(None, description="optional client-provided id (not stored by server)")
|
id: Optional[str] = Field(None, description="optional client id; not stored server-side")
|
||||||
family: Optional[str] = Field(DEFAULT_FAMILY)
|
family: Optional[str] = Field(DEFAULT_FAMILY)
|
||||||
table: Optional[str] = Field(DEFAULT_TABLE)
|
table: Optional[str] = Field(DEFAULT_TABLE)
|
||||||
chain: Optional[str] = Field(DEFAULT_CHAIN)
|
chain: Optional[str] = Field(DEFAULT_CHAIN)
|
||||||
@@ -68,198 +63,162 @@ class ReplaceResult(BaseModel):
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------- nft binding helpers ----------------------
|
# ---------------------- nft binding helpers ----------------------
|
||||||
def _ensure_binding():
|
def _ensure_binding_available():
|
||||||
if Nftables is None:
|
if Nftables is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"python nftables binding not available. Install python3-nftables (system package) or pip-nftables."
|
"python nftables binding not available. Install system package `python3-nftables` or a compatible binding."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def nft_cmd(cmd: str) -> Tuple[int, str, str]:
|
def nft_cmd(cmd: str):
|
||||||
"""Run a libnftables command and return (rc, stdout, stderr)."""
|
"""
|
||||||
_ensure_binding()
|
Run a libnftables command string.
|
||||||
|
Returns tuple (rc, stdout, stderr) as strings. Raises RuntimeError on unexpected binding error.
|
||||||
|
"""
|
||||||
|
_ensure_binding_available()
|
||||||
nft = Nftables()
|
nft = Nftables()
|
||||||
try:
|
try:
|
||||||
rc, out, err = nft.cmd(cmd)
|
rc, out, err = nft.cmd(cmd)
|
||||||
# ensure strings
|
# ensure we work with strings
|
||||||
if isinstance(out, bytes):
|
if isinstance(out, bytes):
|
||||||
out = out.decode(errors="ignore")
|
out = out.decode(errors="ignore")
|
||||||
if isinstance(err, bytes):
|
if isinstance(err, bytes):
|
||||||
err = err.decode(errors="ignore")
|
err = err.decode(errors="ignore")
|
||||||
return rc, out or "", err or ""
|
return rc, out or "", err or ""
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("nft binding call failed for command: %s", cmd)
|
logger.exception("unexpected nft binding exception for cmd=%s", cmd)
|
||||||
raise RuntimeError(f"nft binding failed: {e}")
|
raise RuntimeError(f"nft binding error: {e}")
|
||||||
|
|
||||||
|
|
||||||
def nft_run_or_raise(cmd: str) -> str:
|
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)
|
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))
|
logger.debug("nft cmd: %s -> rc=%s out_len=%d err_len=%d", cmd, rc, len(out), len(err))
|
||||||
if rc != 0:
|
if rc != 0:
|
||||||
raise RuntimeError(err or f"nft command {cmd} failed with rc={rc}")
|
# surface err if available, otherwise a generic message
|
||||||
|
raise RuntimeError(err or f"nft command '{cmd}' failed (rc={rc})")
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
# ---------------------- reconstruction logic ----------------------
|
# ---------------------- chain/table ensure / reconstruction ----------------------
|
||||||
def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
||||||
"""
|
"""
|
||||||
Ensure table and chain exist. Use libnftables commands; fallback to nft -f temp file
|
Ensure the nft table and chain exist. Use libnftables `add` commands first; on failure
|
||||||
if direct add fails. Raises RuntimeError 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)
|
logger.info("ensuring table %s.%s exists", family, table)
|
||||||
|
|
||||||
|
# try add table
|
||||||
try:
|
try:
|
||||||
nft_run_or_raise(f"add table {family} {table}")
|
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:
|
except RuntimeError as e_table:
|
||||||
logger.info("add table failed: %s; trying -f fallback", e_table)
|
logger.info("add table failed: %s. attempting fallback", e_table)
|
||||||
script = f"table {family} {table} {{ }}\n"
|
# fallback: build script and apply the lines individually
|
||||||
tmp = None
|
script_lines = [f"table {family} {table} {{ }}"]
|
||||||
try:
|
try:
|
||||||
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_tbl_create_", suffix=".nft") as tf:
|
for ln in script_lines:
|
||||||
tmp = tf.name
|
nft_run_or_raise(ln)
|
||||||
tf.write(script)
|
logger.debug("created table %s.%s via fallback lines", family, table)
|
||||||
tf.flush()
|
except RuntimeError as e2:
|
||||||
os.fsync(tf.fileno())
|
logger.error("fallback for creating table failed: %s (original: %s)", e2, e_table)
|
||||||
nft_run_or_raise(f"-f {tmp}")
|
raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2
|
||||||
finally:
|
|
||||||
if tmp and os.path.exists(tmp):
|
|
||||||
try:
|
|
||||||
os.remove(tmp)
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
logger.info("ensuring chain %s in table %s", chain, table)
|
# try add chain
|
||||||
|
logger.info("ensuring chain %s in table %s exists", chain, table)
|
||||||
try:
|
try:
|
||||||
nft_run_or_raise(
|
nft_run_or_raise(
|
||||||
f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}'
|
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:
|
except RuntimeError as e_chain:
|
||||||
logger.info("add chain failed: %s; trying -f fallback", e_chain)
|
logger.info("add chain failed: %s. attempting fallback", e_chain)
|
||||||
script = (
|
script_lines = [
|
||||||
f"table {family} {table} {{\n"
|
f"table {family} {table} {{",
|
||||||
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n"
|
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}",
|
||||||
f"}}\n"
|
f"}}",
|
||||||
)
|
]
|
||||||
tmp = None
|
|
||||||
try:
|
try:
|
||||||
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_chain_create_", suffix=".nft") as tf:
|
for ln in script_lines:
|
||||||
tmp = tf.name
|
nft_run_or_raise(ln)
|
||||||
tf.write(script)
|
logger.debug("created chain %s in %s.%s via fallback lines", chain, family, table)
|
||||||
tf.flush()
|
except RuntimeError as e2:
|
||||||
os.fsync(tf.fileno())
|
logger.error("fallback for creating chain failed: %s (original: %s)", e2, e_chain)
|
||||||
nft_run_or_raise(f"-f {tmp}")
|
raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2
|
||||||
finally:
|
|
||||||
if tmp and os.path.exists(tmp):
|
logger.info("table/chain ensured: %s.%s/%s", family, table, chain)
|
||||||
try:
|
|
||||||
os.remove(tmp)
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_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
|
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------- reconstruct rules from nft JSON (best-effort) ----------------------
|
||||||
def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
|
def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Best-effort mapping from nft expression json to our RuleModel-like dict.
|
Heuristic mapping of nft expression JSON to our RuleModel fields.
|
||||||
This is heuristic: nft JSON shapes differ across kernel/libnftables versions.
|
Covers common shapes: meta (iif/oif/length), cmp/payload (ip proto and ports), verdict (action).
|
||||||
We cover common patterns: meta (iif/oif/length), payload/cmp for ip proto and ports, verdict for action.
|
This is best-effort; complex expressions may not map perfectly.
|
||||||
"""
|
"""
|
||||||
match: Dict[str, Any] = {}
|
match: Dict[str, Any] = {}
|
||||||
action: Dict[str, Any] = {}
|
action: Dict[str, Any] = {}
|
||||||
|
|
||||||
# iterate expressions; keep simple heuristics
|
|
||||||
for e in exprs:
|
for e in exprs:
|
||||||
if "meta" in e:
|
if "meta" in e:
|
||||||
m = e["meta"]
|
m = e["meta"]
|
||||||
# keys differ; check for common names
|
|
||||||
key = m.get("key") or m.get("type")
|
key = m.get("key") or m.get("type")
|
||||||
|
v = m.get("v") or m.get("s") or m.get("value")
|
||||||
if key in ("iifname", "iif"):
|
if key in ("iifname", "iif"):
|
||||||
v = m.get("v") or m.get("s") or m.get("value")
|
|
||||||
if v:
|
if v:
|
||||||
match["iif"] = v
|
match["iif"] = v
|
||||||
elif key in ("oifname", "oif"):
|
elif key in ("oifname", "oif"):
|
||||||
v = m.get("v") or m.get("s") or m.get("value")
|
|
||||||
if v:
|
if v:
|
||||||
match["oif"] = v
|
match["oif"] = v
|
||||||
elif key == "length":
|
elif key == "length":
|
||||||
v = m.get("v") or m.get("s") or m.get("value")
|
|
||||||
if v is not None:
|
if v is not None:
|
||||||
match["meta_length"] = v
|
match["meta_length"] = v
|
||||||
elif "payload" in e:
|
elif "payload" in e:
|
||||||
# payload indicates reading bytes of header; usually followed by a 'cmp' comparing to immediate
|
# payload describes read of header bytes; often next 'cmp' compares it
|
||||||
# store payload description to use when we see a cmp
|
# we tag the payload on the expression to help cmp heuristics
|
||||||
# flatten payload into a marker for later cmp detection
|
e["_payload_hint"] = e["payload"]
|
||||||
e_payload = e["payload"]
|
|
||||||
e["_seen_payload"] = e_payload
|
|
||||||
elif "cmp" in e:
|
elif "cmp" in e:
|
||||||
cmp = e["cmp"]
|
cmp = e["cmp"]
|
||||||
# cmp can have 'left'/'right' or 'data' fields
|
|
||||||
left = cmp.get("left")
|
left = cmp.get("left")
|
||||||
right = cmp.get("right")
|
right = cmp.get("right")
|
||||||
# helper to pull immediate numeric value
|
|
||||||
def _extract_immediate(node):
|
def _extract_immediate(node):
|
||||||
if not node:
|
if not node or not isinstance(node, dict):
|
||||||
return None
|
return None
|
||||||
if isinstance(node, dict):
|
for k in ("immediate", "value", "data", "s", "v"):
|
||||||
for k in ("immediate", "value", "data", "s", "v"):
|
if k in node:
|
||||||
if k in node:
|
val = node[k]
|
||||||
val = node[k]
|
if isinstance(val, str) and val.startswith("0x"):
|
||||||
# hex string -> int
|
try:
|
||||||
if isinstance(val, str) and val.startswith("0x"):
|
return int(val, 16)
|
||||||
try:
|
except Exception:
|
||||||
return int(val, 16)
|
return val
|
||||||
except Exception:
|
return val
|
||||||
return val
|
|
||||||
return val
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
imm_left = _extract_immediate(left)
|
imm_left = _extract_immediate(left)
|
||||||
imm_right = _extract_immediate(right)
|
imm_right = _extract_immediate(right)
|
||||||
|
# ip proto numeric likely in 1..255
|
||||||
# If either immediate looks like small numeric, assume ip_proto
|
|
||||||
for imm in (imm_left, imm_right):
|
for imm in (imm_left, imm_right):
|
||||||
if isinstance(imm, int) and 0 < imm < 256:
|
if isinstance(imm, int) and 0 < imm < 256:
|
||||||
# set ip_proto numeric
|
|
||||||
match["ip_proto"] = imm
|
match["ip_proto"] = imm
|
||||||
break
|
break
|
||||||
|
# port heuristics (1..65535)
|
||||||
# If immediate value looks like a TCP/UDP port (typical 1..65535) and payload context indicates tcp/udp dport,
|
|
||||||
# it's hard to be 100% sure; we use heuristic: if cmp mentions 'dport' or payload had offset consistent with port,
|
|
||||||
# or imm is in port range but no ip_proto set, we attempt to set tcp_dport/udp_dport.
|
|
||||||
imm = imm_left if imm_left is not None else imm_right
|
imm = imm_left if imm_left is not None else imm_right
|
||||||
if isinstance(imm, int) and 0 < imm <= 65535:
|
if isinstance(imm, int) and 0 < imm <= 65535:
|
||||||
# if we previously saw payload indicating tcp or udp, detect from payload description
|
# best-effort assign to tcp_dport (common case)
|
||||||
# naive heuristic: if any seen payload mentions 'tcp' or 'udp' in its dict then assign accordingly
|
# more advanced heuristics could inspect payload hints
|
||||||
# otherwise assign tcp_dport by default (best-effort)
|
match.setdefault("tcp_dport", imm)
|
||||||
assigned = False
|
|
||||||
for ev in (left, right):
|
|
||||||
if isinstance(ev, dict):
|
|
||||||
# look for hints
|
|
||||||
if "payload" in ev:
|
|
||||||
pd = ev["payload"]
|
|
||||||
if isinstance(pd, dict) and ("tcp" in str(pd).lower()):
|
|
||||||
match["tcp_dport"] = imm
|
|
||||||
assigned = True
|
|
||||||
break
|
|
||||||
if not assigned:
|
|
||||||
# fallback to tcp_dport heuristic
|
|
||||||
match.setdefault("tcp_dport", imm)
|
|
||||||
elif "verdict" in e:
|
elif "verdict" in e:
|
||||||
v = e["verdict"]
|
v = e["verdict"]
|
||||||
# typical shapes: {"verdict":"accept"} or {"verdict":{"type":"drop"}}
|
# shapes vary: dict or string
|
||||||
if isinstance(v, dict):
|
if isinstance(v, dict):
|
||||||
t = v.get("type") or v.get("kind")
|
t = v.get("type") or v.get("kind")
|
||||||
if t:
|
if t:
|
||||||
action["type"] = t
|
action["type"] = t
|
||||||
# redirect/queue handling may vary; try to extract
|
# redirect / queue extra fields vary
|
||||||
if "to" in v:
|
if "to" in v:
|
||||||
action["type"] = "redirect"
|
action["type"] = "redirect"
|
||||||
action["redirect_port"] = v.get("to")
|
action["redirect_port"] = v.get("to")
|
||||||
@@ -268,25 +227,18 @@ def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
|
|||||||
action["queue_num"] = v.get("queue")
|
action["queue_num"] = v.get("queue")
|
||||||
elif isinstance(v, str):
|
elif isinstance(v, str):
|
||||||
action["type"] = v
|
action["type"] = v
|
||||||
elif "jump" in e:
|
# other expression types intentionally ignored for stateless reconstruction
|
||||||
# jump is effectively a control flow; not modeled
|
|
||||||
pass
|
|
||||||
elif "immediate" in e:
|
|
||||||
# sometimes immediate verdicts
|
|
||||||
pass
|
|
||||||
# there are many other expression types; above covers common ones
|
|
||||||
|
|
||||||
# Default action if none found
|
|
||||||
if "type" not in action:
|
if "type" not in action:
|
||||||
action["type"] = "accept" # kernel often has policy accept if not specified
|
# if kernel default, we assume accept
|
||||||
|
action["type"] = "accept"
|
||||||
return {"match": match, "action": action}
|
return {"match": match, "action": action}
|
||||||
|
|
||||||
|
|
||||||
def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
|
def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Reconstruct a list of rules from the nft kernel listing (JSON).
|
Reconstruct an ordered list of rule dicts from `nft list table` JSON via the binding.
|
||||||
Returns ordered list of dicts each matching RuleModel shape (without id).
|
Returns list of dicts shaped like RuleModel (without id).
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
out = nft_run_or_raise(f"list table {family} {table}")
|
out = nft_run_or_raise(f"list table {family} {table}")
|
||||||
@@ -294,35 +246,33 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
|
|||||||
logger.debug("no table %s.%s found when listing rules", family, table)
|
logger.debug("no table %s.%s found when listing rules", family, table)
|
||||||
return []
|
return []
|
||||||
|
|
||||||
# libnftables returns spammy text sometimes; prefer JSON output
|
# attempt to parse JSON output (binding can return textual JSON)
|
||||||
# attempt to parse as JSON if lib produced JSON; otherwise fallback to textual parsing
|
|
||||||
try:
|
try:
|
||||||
# if out is JSON text produced by libnftables, it is already JSON representation of ruleset
|
|
||||||
data = json.loads(out)
|
data = json.loads(out)
|
||||||
except Exception:
|
except Exception:
|
||||||
# fallback: run with JSON mode by using the binding directly to request JSON output
|
# force JSON mode in binding if previous parsing failed
|
||||||
_ensure_binding()
|
_ensure_binding_available()
|
||||||
nft = Nftables()
|
nft = Nftables()
|
||||||
nft.set_json_output(True)
|
nft.set_json_output(True)
|
||||||
rc, out_json, err = nft.cmd(f"list table {family} {table}")
|
rc, out_json, err = nft.cmd(f"list table {family} {table}")
|
||||||
if rc != 0:
|
if rc != 0:
|
||||||
logger.debug("nft JSON list failed: %s", err)
|
logger.debug("nft JSON listing failed: %s", err)
|
||||||
return []
|
return []
|
||||||
data = json.loads(out_json)
|
data = json.loads(out_json)
|
||||||
|
|
||||||
results: List[Dict[str, Any]] = []
|
results: List[Dict[str, Any]] = []
|
||||||
# nft JSON structure: {"nftables":[ { "table":...}, { "chain":...}, { "rule": {...} }, ... ]}
|
|
||||||
for item in data.get("nftables", []):
|
for item in data.get("nftables", []):
|
||||||
if "rule" not in item:
|
if "rule" not in item:
|
||||||
continue
|
continue
|
||||||
rule_obj = item["rule"]
|
rule_obj = item["rule"]
|
||||||
exprs = rule_obj.get("expr", []) or rule_obj.get("expr", []) # different bindings key names
|
exprs = rule_obj.get("expr", []) or rule_obj.get("expressions", []) or []
|
||||||
recon = _reconstruct_rule_from_exprs(exprs)
|
recon = _reconstruct_rule_from_exprs(exprs)
|
||||||
# create a RuleModel-like dict
|
# the chain name may be in the rule metadata
|
||||||
|
chain_name = rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN
|
||||||
rule_dict = {
|
rule_dict = {
|
||||||
"family": family,
|
"family": family,
|
||||||
"table": table,
|
"table": table,
|
||||||
"chain": rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN,
|
"chain": chain_name,
|
||||||
"match": recon.get("match", {}),
|
"match": recon.get("match", {}),
|
||||||
"action": recon.get("action", {}),
|
"action": recon.get("action", {}),
|
||||||
}
|
}
|
||||||
@@ -335,11 +285,10 @@ def list_rules_from_nft(family: str, table: str) -> List[Dict[str, Any]]:
|
|||||||
@router.get("/rules")
|
@router.get("/rules")
|
||||||
def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
|
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)
|
logger.debug("GET /nft/rules called (family=%s table=%s chain=%s)", family, table, chain)
|
||||||
# ensure chain exists (create if missing) to provide consistent output
|
|
||||||
try:
|
try:
|
||||||
ensure_table_chain(family, table, chain)
|
ensure_table_chain(family, table, chain)
|
||||||
except RuntimeError as e:
|
except RuntimeError as e:
|
||||||
logger.error("failed to ensure table/chain: %s", 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}")
|
raise HTTPException(status_code=500, detail=f"failed to ensure nft table/chain: {e}")
|
||||||
|
|
||||||
rules = list_rules_from_nft(family, table)
|
rules = list_rules_from_nft(family, table)
|
||||||
@@ -356,15 +305,15 @@ def put_rules(
|
|||||||
chain: Optional[str] = DEFAULT_CHAIN,
|
chain: Optional[str] = DEFAULT_CHAIN,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Replace entire ordered rule set. This implementation builds nft commands (without comments)
|
Replace entire ordered ruleset. Validates provided rules, then builds nft commands
|
||||||
from the provided rules and applies via 'nft -f' using libnftables fallback.
|
and applies them line-by-line via libnftables binding (no '-f' token passed into binding).
|
||||||
"""
|
"""
|
||||||
# optimistic concurrency
|
# optimistic concurrency
|
||||||
global _current_version
|
global _current_version
|
||||||
if if_match is not None and _current_version is not None and if_match != _current_version:
|
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")
|
raise HTTPException(status_code=409, detail="version mismatch; fetch latest rules and retry")
|
||||||
|
|
||||||
# validate and ensure table/chain exist
|
# validate per-rule family/table/chain (quick checks)
|
||||||
for r in rules:
|
for r in rules:
|
||||||
if r.family and r.family != family:
|
if r.family and r.family != family:
|
||||||
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}")
|
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}")
|
||||||
@@ -373,21 +322,24 @@ def put_rules(
|
|||||||
if r.chain and r.chain != chain:
|
if r.chain and r.chain != chain:
|
||||||
raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}")
|
raise HTTPException(status_code=400, detail=f"rule chain mismatch: {r.chain} != {chain}")
|
||||||
|
|
||||||
|
# ensure table/chain exist first
|
||||||
try:
|
try:
|
||||||
ensure_table_chain(family, table, chain)
|
ensure_table_chain(family, table, chain)
|
||||||
except RuntimeError as e:
|
except RuntimeError as e:
|
||||||
logger.error("failed to ensure table/chain: %s", e)
|
logger.error("failed to ensure table/chain: %s", e)
|
||||||
raise HTTPException(status_code=500, detail=str(e))
|
raise HTTPException(status_code=500, detail=str(e))
|
||||||
|
|
||||||
# Build nft script lines (no comments)
|
# build script lines (flush + ordered add rules)
|
||||||
script_lines: List[str] = []
|
script_lines: List[str] = []
|
||||||
script_lines.append(f"flush chain {family} {table} {chain}")
|
script_lines.append(f"flush chain {family} {table} {chain}")
|
||||||
|
|
||||||
for r in rules:
|
for r in rules:
|
||||||
# ensure id exists only for client convenience; not stored
|
# assign client-side id if missing (not stored)
|
||||||
if not r.id:
|
if not r.id:
|
||||||
r.id = str(uuid.uuid4())
|
r.id = str(uuid.uuid4())
|
||||||
# build add rule command using the same helpers
|
|
||||||
match_frag = []
|
# build match fragments
|
||||||
|
match_frag: List[str] = []
|
||||||
if r.match.iif:
|
if r.match.iif:
|
||||||
match_frag += ["iif", f'"{r.match.iif}"']
|
match_frag += ["iif", f'"{r.match.iif}"']
|
||||||
if r.match.oif:
|
if r.match.oif:
|
||||||
@@ -401,43 +353,48 @@ def put_rules(
|
|||||||
if r.match.udp_dport:
|
if r.match.udp_dport:
|
||||||
match_frag += ["udp", "dport", str(r.match.udp_dport)]
|
match_frag += ["udp", "dport", str(r.match.udp_dport)]
|
||||||
|
|
||||||
action_frag = []
|
# build action fragment
|
||||||
if r.action.type == "drop":
|
action_frag: List[str] = []
|
||||||
|
a_type = r.action.type or "accept"
|
||||||
|
if a_type == "drop":
|
||||||
action_frag = ["drop"]
|
action_frag = ["drop"]
|
||||||
elif r.action.type == "accept" or not r.action.type:
|
elif a_type == "accept":
|
||||||
action_frag = ["accept"]
|
action_frag = ["accept"]
|
||||||
elif r.action.type == "queue":
|
elif a_type == "queue":
|
||||||
num = r.action.queue_num if r.action.queue_num is not None else 0
|
num = r.action.queue_num if r.action.queue_num is not None else 0
|
||||||
action_frag = ["queue", "num", str(num)]
|
action_frag = ["queue", "num", str(num)]
|
||||||
elif r.action.type == "redirect":
|
elif a_type == "redirect":
|
||||||
if r.action.redirect_port is None:
|
if r.action.redirect_port is None:
|
||||||
raise HTTPException(status_code=400, detail="redirect action requires redirect_port")
|
raise HTTPException(status_code=400, detail="redirect action requires redirect_port")
|
||||||
action_frag = ["redirect", "to", f":{r.action.redirect_port}"]
|
action_frag = ["redirect", "to", f":{r.action.redirect_port}"]
|
||||||
else:
|
else:
|
||||||
raise HTTPException(status_code=400, detail=f"unsupported action type: {r.action.type}")
|
raise HTTPException(status_code=400, detail=f"unsupported action type: {a_type}")
|
||||||
|
|
||||||
parts = ["add", "rule", r.family, r.table, r.chain]
|
parts = ["add", "rule", r.family, r.table, r.chain] + match_frag + action_frag
|
||||||
parts += match_frag
|
|
||||||
parts += action_frag
|
|
||||||
script_lines.append(" ".join(parts))
|
script_lines.append(" ".join(parts))
|
||||||
|
|
||||||
script = "\n".join(script_lines) + "\n"
|
# apply script lines via binding (line-by-line)
|
||||||
tmpfile_path: Optional[str] = None
|
tmpfile_path: Optional[str] = None
|
||||||
try:
|
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:
|
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf:
|
||||||
tmpfile_path = tf.name
|
tmpfile_path = tf.name
|
||||||
tf.write(script)
|
tf.write("\n".join(script_lines) + "\n")
|
||||||
tf.flush()
|
tf.flush()
|
||||||
os.fsync(tf.fileno())
|
os.fsync(tf.fileno())
|
||||||
logger.info("wrote nft script to %s; applying...", tmpfile_path)
|
logger.info("wrote nft script to %s; applying line-by-line via binding...", tmpfile_path)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# use libnftables binding to apply file if possible, otherwise lib will call nft -f underneath
|
for ln in script_lines:
|
||||||
nft_run_or_raise(f"-f {tmpfile_path}")
|
ln = ln.strip()
|
||||||
|
if not ln:
|
||||||
|
continue
|
||||||
|
nft_run_or_raise(ln)
|
||||||
except RuntimeError as e:
|
except RuntimeError as e:
|
||||||
logger.error("failed applying nft script: %s", 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}")
|
raise HTTPException(status_code=500, detail=f"failed applying nft script: {e}")
|
||||||
|
|
||||||
|
# success -> bump version
|
||||||
_current_version = str(uuid.uuid4())
|
_current_version = str(uuid.uuid4())
|
||||||
logger.info("applied nft ruleset successfully; version=%s", _current_version)
|
logger.info("applied nft ruleset successfully; version=%s", _current_version)
|
||||||
return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))
|
return ReplaceResult(version=_current_version, applied=True, rules_count=len(rules))
|
||||||
|
|||||||
Reference in New Issue
Block a user