add libnft, simpler api
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s
This commit is contained in:
Binary file not shown.
@@ -1,106 +1,116 @@
|
|||||||
# fastapi_nft_replace.py
|
# fastapi_nft_stateless.py
|
||||||
"""
|
"""
|
||||||
nftables router for FastAPI to manage nftables bridge rules dynamically.
|
Stateless FastAPI nftables router.
|
||||||
|
Only uses kernel-stored info (nft) as the source of truth.
|
||||||
This version will attempt to create any kernel/runtime prerequisites:
|
|
||||||
- load 'bridge' and 'br_netfilter' modules (via modprobe)
|
|
||||||
- enable sysctls net.bridge.bridge-nf-call-iptables and net.bridge.bridge-nf-call-ip6tables
|
|
||||||
- create the nft table and chain (with nft add ... and nft -f fallback)
|
|
||||||
|
|
||||||
Endpoints (mounted under /nft):
|
Endpoints (mounted under /nft):
|
||||||
- GET /rules -> list active rules (read from nftables)
|
- GET /rules -> reconstruct rule objects from nft kernel state (best-effort)
|
||||||
- POST /rules -> add one or many rules (append)
|
- PUT /rules -> replace entire ordered ruleset (applies via nft -f)
|
||||||
- DELETE /rules -> delete one or many rules by id
|
|
||||||
- PUT /rules -> replace entire ordered rule set via nft -f (returns new version)
|
No persistence, no comments used for mapping.
|
||||||
"""
|
"""
|
||||||
from fastapi import APIRouter, HTTPException
|
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, Union, Tuple
|
||||||
import subprocess
|
|
||||||
import uuid
|
import uuid
|
||||||
import json
|
import json
|
||||||
import base64
|
|
||||||
import re
|
|
||||||
import tempfile
|
import tempfile
|
||||||
import os
|
import os
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
# Router and logger ---------------------------------------------------------
|
# libnftables binding
|
||||||
router = APIRouter()
|
try:
|
||||||
|
from nftables import Nftables
|
||||||
|
except Exception:
|
||||||
|
Nftables = None # will raise when used
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/nft", tags=["nftables"])
|
||||||
logger = logging.getLogger("nftables")
|
logger = logging.getLogger("nftables")
|
||||||
logger.debug("nftables router module loaded")
|
logger.debug("nftables stateless router loaded")
|
||||||
|
|
||||||
# Defaults and in-memory version token --------------------------------------
|
# Defaults
|
||||||
DEFAULT_TABLE = "mitm_tbl"
|
DEFAULT_TABLE = "mitm_tbl"
|
||||||
DEFAULT_CHAIN = "forward"
|
DEFAULT_CHAIN = "forward"
|
||||||
DEFAULT_FAMILY = "bridge"
|
DEFAULT_FAMILY = "bridge"
|
||||||
|
|
||||||
# in-memory version token updated on successful PUT
|
|
||||||
_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 # allow integers or ranges like "100-200"
|
meta_length: Optional[Union[int, str]] = None
|
||||||
ip_proto: Optional[str] = None
|
ip_proto: Optional[Union[int, str]] = None
|
||||||
tcp_dport: Optional[int] = None
|
tcp_dport: Optional[int] = None
|
||||||
udp_dport: Optional[int] = None
|
udp_dport: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
class ActionModel(BaseModel):
|
class ActionModel(BaseModel):
|
||||||
type: str # drop | accept | queue | redirect
|
type: Optional[str] = None # drop | accept | queue | redirect
|
||||||
queue_num: Optional[int] = None
|
queue_num: Optional[int] = None
|
||||||
redirect_port: Optional[int] = None
|
redirect_port: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
class RuleModel(BaseModel):
|
class RuleModel(BaseModel):
|
||||||
id: Optional[str] = Field(None, description="optional rule id; generated if missing")
|
# note: since we don't persist IDs, id remains optional
|
||||||
|
id: Optional[str] = Field(None, description="optional client-provided id (not stored by server)")
|
||||||
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)
|
||||||
match: MatchModel
|
match: MatchModel
|
||||||
action: ActionModel
|
action: ActionModel
|
||||||
|
|
||||||
|
|
||||||
class ReplaceResult(BaseModel):
|
class ReplaceResult(BaseModel):
|
||||||
version: str
|
version: str
|
||||||
applied: bool
|
applied: bool
|
||||||
rules_count: int
|
rules_count: int
|
||||||
|
|
||||||
|
|
||||||
# ---------------------- Utilities ----------------------
|
# ---------------------- nft binding helpers ----------------------
|
||||||
def run_nft(args: List[str]) -> Tuple[str, str]:
|
def _ensure_binding():
|
||||||
"""
|
if Nftables is None:
|
||||||
Run nft with given args.
|
raise RuntimeError(
|
||||||
Returns (stdout, stderr).
|
"python nftables binding not available. Install python3-nftables (system package) or pip-nftables."
|
||||||
Raises RuntimeError on non-zero exit with stderr included.
|
)
|
||||||
"""
|
|
||||||
logger.debug("running nft: %s", " ".join(["nft"] + args))
|
|
||||||
try:
|
|
||||||
proc = subprocess.run(["nft"] + args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True)
|
|
||||||
logger.debug("nft stdout: %s", proc.stdout.strip())
|
|
||||||
return proc.stdout, proc.stderr
|
|
||||||
except subprocess.CalledProcessError as e:
|
|
||||||
logger.error("nft failed: %s -- %s", " ".join(e.cmd), e.stderr.strip())
|
|
||||||
raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}")
|
|
||||||
|
|
||||||
|
|
||||||
|
def nft_cmd(cmd: str) -> Tuple[int, str, str]:
|
||||||
|
"""Run a libnftables command and return (rc, stdout, stderr)."""
|
||||||
|
_ensure_binding()
|
||||||
|
nft = Nftables()
|
||||||
|
try:
|
||||||
|
rc, out, err = nft.cmd(cmd)
|
||||||
|
# ensure 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("nft binding call failed for command: %s", cmd)
|
||||||
|
raise RuntimeError(f"nft binding failed: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
def nft_run_or_raise(cmd: str) -> str:
|
||||||
|
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:
|
||||||
|
raise RuntimeError(err or f"nft command {cmd} failed with rc={rc}")
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------- reconstruction logic ----------------------
|
||||||
def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
||||||
"""
|
"""
|
||||||
Ensure the nft table and chain exist. On serious failures this raises RuntimeError.
|
Ensure table and chain exist. Use libnftables commands; fallback to nft -f temp file
|
||||||
This function:
|
if direct add fails. Raises RuntimeError on failure.
|
||||||
- attempts `nft add table` and `nft add chain`
|
|
||||||
- if those fail, tries to apply a tiny nft script with `nft -f` to create table and chain
|
|
||||||
"""
|
"""
|
||||||
logger.info("ensuring table %s.%s exists", family, table)
|
logger.info("ensuring table %s.%s exists", family, table)
|
||||||
|
|
||||||
# Try simple add table first
|
|
||||||
try:
|
try:
|
||||||
run_nft(["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("nft add table failed for %s.%s; trying nft -f fallback: %s", family, table, e_table)
|
logger.info("add table failed: %s; trying -f fallback", e_table)
|
||||||
# fallback create via nft -f script
|
|
||||||
script = f"table {family} {table} {{ }}\n"
|
script = f"table {family} {table} {{ }}\n"
|
||||||
tmp = None
|
tmp = None
|
||||||
try:
|
try:
|
||||||
@@ -109,12 +119,7 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
|||||||
tf.write(script)
|
tf.write(script)
|
||||||
tf.flush()
|
tf.flush()
|
||||||
os.fsync(tf.fileno())
|
os.fsync(tf.fileno())
|
||||||
try:
|
nft_run_or_raise(f"-f {tmp}")
|
||||||
run_nft(["-f", tmp])
|
|
||||||
logger.debug("created table %s.%s via nft -f", family, table)
|
|
||||||
except RuntimeError as e2:
|
|
||||||
logger.error("nft -f fallback to create table failed: %s (original: %s)", e2, e_table)
|
|
||||||
raise RuntimeError(f"failed to create nft table {family}.{table}: {e2}") from e2
|
|
||||||
finally:
|
finally:
|
||||||
if tmp and os.path.exists(tmp):
|
if tmp and os.path.exists(tmp):
|
||||||
try:
|
try:
|
||||||
@@ -122,13 +127,13 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# Now ensure chain exists
|
logger.info("ensuring chain %s in table %s", chain, table)
|
||||||
logger.info("ensuring chain %s in table %s exists", chain, table)
|
|
||||||
try:
|
try:
|
||||||
run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"])
|
nft_run_or_raise(
|
||||||
logger.debug("created chain %s in %s.%s via add chain", chain, family, table)
|
f'add chain {family} {table} {chain} {{ type filter hook forward priority 0; policy accept; }}'
|
||||||
|
)
|
||||||
except RuntimeError as e_chain:
|
except RuntimeError as e_chain:
|
||||||
logger.info("nft add chain failed for %s in %s.%s; trying nft -f fallback: %s", chain, family, table, e_chain)
|
logger.info("add chain failed: %s; trying -f fallback", e_chain)
|
||||||
script = (
|
script = (
|
||||||
f"table {family} {table} {{\n"
|
f"table {family} {table} {{\n"
|
||||||
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n"
|
f" chain {chain} {{ type filter hook forward priority 0; policy accept; }}\n"
|
||||||
@@ -141,12 +146,7 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
|||||||
tf.write(script)
|
tf.write(script)
|
||||||
tf.flush()
|
tf.flush()
|
||||||
os.fsync(tf.fileno())
|
os.fsync(tf.fileno())
|
||||||
try:
|
nft_run_or_raise(f"-f {tmp}")
|
||||||
run_nft(["-f", tmp])
|
|
||||||
logger.debug("created chain %s in %s.%s via nft -f", chain, family, table)
|
|
||||||
except RuntimeError as e2:
|
|
||||||
logger.error("nft -f fallback to create chain failed: %s (original: %s)", e2, e_chain)
|
|
||||||
raise RuntimeError(f"failed to create nft chain {chain} in {family}.{table}: {e2}") from e2
|
|
||||||
finally:
|
finally:
|
||||||
if tmp and os.path.exists(tmp):
|
if tmp and os.path.exists(tmp):
|
||||||
try:
|
try:
|
||||||
@@ -154,210 +154,217 @@ def ensure_table_chain(family: str, table: str, chain: str) -> None:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
logger.info("table/chain ensured: %s.%s/%s", family, table, chain)
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
def encode_rule_comment(rule: Dict[str, Any]) -> str:
|
def _reconstruct_rule_from_exprs(exprs: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||||
# store rule JSON as base64 to avoid quoting/escaping issues inside nft comment
|
"""
|
||||||
j = json.dumps(rule, separators=(",", ":"))
|
Best-effort mapping from nft expression json to our RuleModel-like dict.
|
||||||
b = base64.b64encode(j.encode()).decode()
|
This is heuristic: nft JSON shapes differ across kernel/libnftables versions.
|
||||||
rid = rule.get("id") or ""
|
We cover common patterns: meta (iif/oif/length), payload/cmp for ip proto and ports, verdict for action.
|
||||||
return f"mitm_id:{rid} mitm_json:{b}"
|
"""
|
||||||
|
match: Dict[str, Any] = {}
|
||||||
|
action: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
# iterate expressions; keep simple heuristics
|
||||||
def decode_comment_payload(comment: str) -> Optional[Dict[str, Any]]:
|
for e in exprs:
|
||||||
# expects comment like: mitm_id:<id> mitm_json:<base64>
|
if "meta" in e:
|
||||||
|
m = e["meta"]
|
||||||
|
# keys differ; check for common names
|
||||||
|
key = m.get("key") or m.get("type")
|
||||||
|
if key in ("iifname", "iif"):
|
||||||
|
v = m.get("v") or m.get("s") or m.get("value")
|
||||||
|
if v:
|
||||||
|
match["iif"] = v
|
||||||
|
elif key in ("oifname", "oif"):
|
||||||
|
v = m.get("v") or m.get("s") or m.get("value")
|
||||||
|
if v:
|
||||||
|
match["oif"] = v
|
||||||
|
elif key == "length":
|
||||||
|
v = m.get("v") or m.get("s") or m.get("value")
|
||||||
|
if v is not None:
|
||||||
|
match["meta_length"] = v
|
||||||
|
elif "payload" in e:
|
||||||
|
# payload indicates reading bytes of header; usually followed by a 'cmp' comparing to immediate
|
||||||
|
# store payload description to use when we see a cmp
|
||||||
|
# flatten payload into a marker for later cmp detection
|
||||||
|
e_payload = e["payload"]
|
||||||
|
e["_seen_payload"] = e_payload
|
||||||
|
elif "cmp" in e:
|
||||||
|
cmp = e["cmp"]
|
||||||
|
# cmp can have 'left'/'right' or 'data' fields
|
||||||
|
left = cmp.get("left")
|
||||||
|
right = cmp.get("right")
|
||||||
|
# helper to pull immediate numeric value
|
||||||
|
def _extract_immediate(node):
|
||||||
|
if not node:
|
||||||
|
return None
|
||||||
|
if isinstance(node, dict):
|
||||||
|
for k in ("immediate", "value", "data", "s", "v"):
|
||||||
|
if k in node:
|
||||||
|
val = node[k]
|
||||||
|
# hex string -> int
|
||||||
|
if isinstance(val, str) and val.startswith("0x"):
|
||||||
try:
|
try:
|
||||||
parts = comment.split()
|
return int(val, 16)
|
||||||
kv = {p.split(":", 1)[0]: p.split(":", 1)[1] for p in parts if ":" in p}
|
|
||||||
b64 = kv.get("mitm_json")
|
|
||||||
if not b64:
|
|
||||||
return None
|
|
||||||
j = base64.b64decode(b64.encode()).decode()
|
|
||||||
return json.loads(j)
|
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.debug("failed to decode comment payload: %s", comment)
|
return val
|
||||||
|
return val
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
imm_left = _extract_immediate(left)
|
||||||
|
imm_right = _extract_immediate(right)
|
||||||
|
|
||||||
def build_nft_match_fragment(match: MatchModel) -> List[str]:
|
# If either immediate looks like small numeric, assume ip_proto
|
||||||
frag: List[str] = []
|
for imm in (imm_left, imm_right):
|
||||||
if match.iif:
|
if isinstance(imm, int) and 0 < imm < 256:
|
||||||
frag += ["iif", f'"{match.iif}"']
|
# set ip_proto numeric
|
||||||
if match.oif:
|
match["ip_proto"] = imm
|
||||||
frag += ["oif", f'"{match.oif}"']
|
break
|
||||||
if match.meta_length is not None:
|
|
||||||
frag += ["meta", "length", str(match.meta_length)]
|
|
||||||
if match.ip_proto:
|
|
||||||
frag += ["ip", "protocol", 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
|
|
||||||
|
|
||||||
|
# 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
|
||||||
|
if isinstance(imm, int) and 0 < imm <= 65535:
|
||||||
|
# if we previously saw payload indicating tcp or udp, detect from payload description
|
||||||
|
# naive heuristic: if any seen payload mentions 'tcp' or 'udp' in its dict then assign accordingly
|
||||||
|
# otherwise assign tcp_dport by default (best-effort)
|
||||||
|
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:
|
||||||
|
v = e["verdict"]
|
||||||
|
# typical shapes: {"verdict":"accept"} or {"verdict":{"type":"drop"}}
|
||||||
|
if isinstance(v, dict):
|
||||||
|
t = v.get("type") or v.get("kind")
|
||||||
|
if t:
|
||||||
|
action["type"] = t
|
||||||
|
# redirect/queue handling may vary; try to extract
|
||||||
|
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
|
||||||
|
elif "jump" in e:
|
||||||
|
# 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
|
||||||
|
|
||||||
def build_nft_action_fragment(action: ActionModel) -> List[str]:
|
# Default action if none found
|
||||||
if action.type == "drop":
|
if "type" not in action:
|
||||||
return ["drop"]
|
action["type"] = "accept" # kernel often has policy accept if not specified
|
||||||
if action.type == "accept":
|
|
||||||
return ["accept"]
|
|
||||||
if action.type == "queue":
|
|
||||||
num = action.queue_num if action.queue_num is not None else 0
|
|
||||||
return ["queue", "num", str(num)]
|
|
||||||
if action.type == "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}")
|
|
||||||
|
|
||||||
|
return {"match": match, "action": action}
|
||||||
def nft_rule_line_from_model(rule: RuleModel) -> str:
|
|
||||||
"""
|
|
||||||
Produce a single-line nft command:
|
|
||||||
add rule <family> <table> <chain> <match...> <action...> comment "<encoded>"
|
|
||||||
"""
|
|
||||||
match_frag = build_nft_match_fragment(rule.match)
|
|
||||||
action_frag = build_nft_action_fragment(rule.action)
|
|
||||||
comment = encode_rule_comment(rule.dict())
|
|
||||||
|
|
||||||
parts = ["add", "rule", rule.family, rule.table, rule.chain]
|
|
||||||
parts += match_frag
|
|
||||||
parts += action_frag
|
|
||||||
return " ".join(parts) + f' comment "{comment}"'
|
|
||||||
|
|
||||||
|
|
||||||
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).
|
||||||
|
Returns ordered list of dicts each matching RuleModel shape (without id).
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
out, _ = run_nft(["list", "table", family, table])
|
out = nft_run_or_raise(f"list table {family} {table}")
|
||||||
except RuntimeError:
|
except RuntimeError:
|
||||||
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 []
|
||||||
|
|
||||||
results: List[Dict[str, Any]] = []
|
# libnftables returns spammy text sometimes; prefer JSON output
|
||||||
for line in out.splitlines():
|
# attempt to parse as JSON if lib produced JSON; otherwise fallback to textual parsing
|
||||||
line = line.strip()
|
|
||||||
if "comment" in line and "mitm_json:" in line:
|
|
||||||
try:
|
try:
|
||||||
first_quote = line.index('"')
|
# if out is JSON text produced by libnftables, it is already JSON representation of ruleset
|
||||||
last_quote = line.rindex('"')
|
data = json.loads(out)
|
||||||
comment_str = line[first_quote + 1:last_quote]
|
except Exception:
|
||||||
except ValueError:
|
# fallback: run with JSON mode by using the binding directly to request JSON output
|
||||||
comment_str = line.split("comment", 1)[1].strip()
|
_ensure_binding()
|
||||||
|
nft = Nftables()
|
||||||
parsed = decode_comment_payload(comment_str)
|
nft.set_json_output(True)
|
||||||
if parsed is not None:
|
rc, out_json, err = nft.cmd(f"list table {family} {table}")
|
||||||
results.append(parsed)
|
if rc != 0:
|
||||||
logger.debug("listed %d nft rules from %s.%s", len(results), family, table)
|
logger.debug("nft JSON list failed: %s", err)
|
||||||
return results
|
|
||||||
|
|
||||||
|
|
||||||
def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]:
|
|
||||||
ensure_table_chain(rule.family, rule.table, rule.chain)
|
|
||||||
if not rule.id:
|
|
||||||
rule.id = str(uuid.uuid4())
|
|
||||||
|
|
||||||
cmd_text = nft_rule_line_from_model(rule)
|
|
||||||
try:
|
|
||||||
# run as discrete args to avoid shell quoting issues
|
|
||||||
run_nft(cmd_text.split())
|
|
||||||
logger.info("added nft rule id=%s family=%s table=%s chain=%s", rule.id, rule.family, rule.table, rule.chain)
|
|
||||||
return {"id": rule.id, "status": "added"}
|
|
||||||
except RuntimeError as e:
|
|
||||||
logger.error("failed to add rule id=%s: %s", rule.id, e)
|
|
||||||
raise HTTPException(status_code=500, detail=str(e))
|
|
||||||
|
|
||||||
|
|
||||||
def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]:
|
|
||||||
try:
|
|
||||||
out, _ = run_nft(["list", "chain", family, table, chain, "-a"])
|
|
||||||
except RuntimeError:
|
|
||||||
logger.debug("no chain %s in table %s.%s when attempting delete", chain, family, table)
|
|
||||||
return []
|
return []
|
||||||
|
data = json.loads(out_json)
|
||||||
|
|
||||||
deleted: List[str] = []
|
results: List[Dict[str, Any]] = []
|
||||||
for line in out.splitlines():
|
# nft JSON structure: {"nftables":[ { "table":...}, { "chain":...}, { "rule": {...} }, ... ]}
|
||||||
if "comment" in line and "mitm_id:" in line:
|
for item in data.get("nftables", []):
|
||||||
try:
|
if "rule" not in item:
|
||||||
q1 = line.index('"')
|
|
||||||
q2 = line.index('"', q1 + 1)
|
|
||||||
comment_str = line[q1 + 1:q2]
|
|
||||||
except ValueError:
|
|
||||||
comment_str = line.split("comment", 1)[1]
|
|
||||||
|
|
||||||
parsed = decode_comment_payload(comment_str)
|
|
||||||
if not parsed:
|
|
||||||
continue
|
continue
|
||||||
rid = parsed.get("id")
|
rule_obj = item["rule"]
|
||||||
if rid in ids:
|
exprs = rule_obj.get("expr", []) or rule_obj.get("expr", []) # different bindings key names
|
||||||
m = re.search(r"handle\s+(\d+)", line)
|
recon = _reconstruct_rule_from_exprs(exprs)
|
||||||
if not m:
|
# create a RuleModel-like dict
|
||||||
parts = line.split()
|
rule_dict = {
|
||||||
if "handle" in parts:
|
"family": family,
|
||||||
hi = parts.index("handle")
|
"table": table,
|
||||||
if hi + 1 < len(parts):
|
"chain": rule_obj.get("chain") or rule_obj.get("chain_name") or DEFAULT_CHAIN,
|
||||||
handle = parts[hi + 1]
|
"match": recon.get("match", {}),
|
||||||
else:
|
"action": recon.get("action", {}),
|
||||||
continue
|
}
|
||||||
else:
|
results.append(rule_dict)
|
||||||
continue
|
logger.debug("reconstructed %d rules from %s.%s", len(results), family, table)
|
||||||
else:
|
return results
|
||||||
handle = m.group(1)
|
|
||||||
|
|
||||||
try:
|
|
||||||
run_nft(["delete", "rule", family, table, chain, "handle", handle])
|
|
||||||
logger.info("deleted nft rule id=%s handle=%s", rid, handle)
|
|
||||||
deleted.append(rid)
|
|
||||||
except RuntimeError:
|
|
||||||
logger.error("failed to delete nft rule id=%s handle=%s", rid, handle)
|
|
||||||
continue
|
|
||||||
return deleted
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------- API endpoints ----------------------
|
# ---------------------- API endpoints ----------------------
|
||||||
@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):
|
||||||
"""
|
|
||||||
Ensure the table/chain exist (try to create them if missing), then return
|
|
||||||
the rules from nftables. If ensure_table_chain fails we return 500 with a clear message.
|
|
||||||
"""
|
|
||||||
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 required table/chain exist before listing rules
|
|
||||||
ensure_table_chain(family, table, chain)
|
ensure_table_chain(family, table, chain)
|
||||||
except RuntimeError as e:
|
except RuntimeError as e:
|
||||||
logger.error("failed to create/ensure nft table/chain %s.%s/%s: %s", family, table, chain, e)
|
logger.error("failed to ensure table/chain: %s", e)
|
||||||
# return an HTTP 500 so the frontend knows setup failed
|
|
||||||
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}")
|
||||||
|
|
||||||
# now safe to list rules
|
|
||||||
rules = list_rules_from_nft(family, table)
|
rules = list_rules_from_nft(family, table)
|
||||||
return {"count": len(rules), "rules": rules, "version": _current_version}
|
return {"count": len(rules), "rules": rules, "version": _current_version}
|
||||||
|
|
||||||
|
|
||||||
@router.post("/rules")
|
|
||||||
def post_rules(payload: Union[RuleModel, List[RuleModel]]):
|
|
||||||
rules = payload if isinstance(payload, list) else [payload]
|
|
||||||
results = []
|
|
||||||
for r in rules:
|
|
||||||
res = add_rule_to_nft(r)
|
|
||||||
results.append(res)
|
|
||||||
return {"results": results}
|
|
||||||
|
|
||||||
|
|
||||||
@router.delete("/rules")
|
|
||||||
def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT_FAMILY,
|
|
||||||
table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
|
|
||||||
ids = [payload] if isinstance(payload, str) else payload
|
|
||||||
deleted = delete_rules_by_ids(ids, family, table, chain)
|
|
||||||
results = [{"id": i, "deleted": i in deleted} for i in ids]
|
|
||||||
return {"results": results}
|
|
||||||
|
|
||||||
|
|
||||||
@router.put("/rules", response_model=ReplaceResult)
|
@router.put("/rules", response_model=ReplaceResult)
|
||||||
def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY,
|
def put_rules(
|
||||||
table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
|
rules: List[RuleModel],
|
||||||
# validate per-rule family/table/chain if present
|
request: Request,
|
||||||
|
if_match: Optional[str] = Header(None, alias="If-Match"),
|
||||||
|
family: Optional[str] = DEFAULT_FAMILY,
|
||||||
|
table: Optional[str] = DEFAULT_TABLE,
|
||||||
|
chain: Optional[str] = DEFAULT_CHAIN,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Replace entire ordered rule set. This implementation builds nft commands (without comments)
|
||||||
|
from the provided rules and applies via 'nft -f' using libnftables fallback.
|
||||||
|
"""
|
||||||
|
# optimistic concurrency
|
||||||
|
global _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")
|
||||||
|
|
||||||
|
# validate and ensure table/chain exist
|
||||||
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}")
|
||||||
@@ -366,32 +373,55 @@ def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY,
|
|||||||
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 (this will attempt to modprobe + sysctl if needed)
|
|
||||||
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))
|
||||||
|
|
||||||
# ensure rule IDs
|
# Build nft script lines (no comments)
|
||||||
|
script_lines: List[str] = []
|
||||||
|
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
|
||||||
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 = []
|
||||||
|
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 nft script lines
|
action_frag = []
|
||||||
lines: List[str] = []
|
if r.action.type == "drop":
|
||||||
# flush chain (clear existing ordered rules)
|
action_frag = ["drop"]
|
||||||
lines.append(f"flush chain {family} {table} {chain}")
|
elif r.action.type == "accept" or not r.action.type:
|
||||||
|
action_frag = ["accept"]
|
||||||
|
elif r.action.type == "queue":
|
||||||
|
num = r.action.queue_num if r.action.queue_num is not None else 0
|
||||||
|
action_frag = ["queue", "num", str(num)]
|
||||||
|
elif r.action.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: {r.action.type}")
|
||||||
|
|
||||||
# add ordered rules
|
parts = ["add", "rule", r.family, r.table, r.chain]
|
||||||
for r in rules:
|
parts += match_frag
|
||||||
try:
|
parts += action_frag
|
||||||
lines.append(nft_rule_line_from_model(r))
|
script_lines.append(" ".join(parts))
|
||||||
except Exception as e:
|
|
||||||
raise HTTPException(status_code=400, detail=f"invalid rule: {e}")
|
|
||||||
|
|
||||||
script = "\n".join(lines) + "\n"
|
|
||||||
|
|
||||||
|
script = "\n".join(script_lines) + "\n"
|
||||||
tmpfile_path: Optional[str] = None
|
tmpfile_path: Optional[str] = None
|
||||||
try:
|
try:
|
||||||
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:
|
||||||
@@ -401,15 +431,13 @@ def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY,
|
|||||||
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...", tmpfile_path)
|
||||||
|
|
||||||
# apply the file
|
|
||||||
try:
|
try:
|
||||||
run_nft(["-f", tmpfile_path])
|
# use libnftables binding to apply file if possible, otherwise lib will call nft -f underneath
|
||||||
|
nft_run_or_raise(f"-f {tmpfile_path}")
|
||||||
except RuntimeError as e:
|
except RuntimeError as e:
|
||||||
logger.error("failed applying nft script: %s", e)
|
logger.error("failed applying nft script: %s", 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 token
|
|
||||||
global _current_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