add logging and improve api for nftables
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s

This commit is contained in:
2026-01-10 18:15:52 +01:00
parent 04d644365a
commit 0ff884831c

View File

@@ -1,52 +1,44 @@
# fastapi_nft_replace.py
""" """
FastAPI app to manage nftables bridge rules dynamically. nftables router for FastAPI to manage nftables bridge rules dynamically.
- Removed idempotent-by-id behavior: POST /rules always adds rules (generates id if missing)
- Added DELETE /rules endpoint to delete one or many rules by id Endpoints (mounted under /nft):
- GET /rules -> list active rules (read from nftables)
- POST /rules -> add one or many rules (append)
- DELETE /rules -> delete one or many rules by id
- PUT /rules -> replace entire ordered rule set via nft -f (returns new version)
Notes: Notes:
- Runs nft(8) commands; the process must have sufficient privileges (run as root or via sudo). - Must run with privileges to call `nft` (root or via sudo).
- This implementation stores the full rule JSON inside the nft rule comment as base64 to allow round-trip parsing. - Rules are stored in the nft rule comment as base64-encoded JSON for round-trip parsing.
- The API keeps sniffer separate; this module only manages kernel rules. - PUT /rules writes a temporary nft script and runs `nft -f <file>`; this applies the new ordered rules.
Endpoints:
- GET /rules -> list active rules (read from nftables)
- POST /rules -> add one or many rules
- DELETE /rules -> delete one or many rules by id
Rule schema (example):
{
"id": "optional-uuid-if-you-want",
"table": "mitm_tbl",
"chain": "forward",
"family": "bridge",
"match": {
"iif": "br0",
"oif": "eth1",
"meta_length": "100-200", # or single int as str/number
"ip_proto": "tcp",
"tcp_dport": 80
},
"action": {"type": "drop"}
}
Supported matches in this example: iif, oif, meta_length, ip_proto, tcp_dport, udp_dport
Supported actions: drop, accept, queue (num), redirect (port)
""" """
from fastapi import APIRouter, HTTPException
from fastapi import APIRouter, APIRouter, HTTPException
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any, Union from typing import Optional, List, Dict, Any, Union, Tuple
import subprocess import subprocess
import uuid import uuid
import json import json
import base64 import base64
import re import re
import tempfile
import os
import logging
# Router and logger ---------------------------------------------------------
router = APIRouter()
logger = logging.getLogger("nftables")
logger.debug("nftables router module loaded")
# Defaults and in-memory version token --------------------------------------
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
# ---------------------- Pydantic models ---------------------- # ---------------------- Pydantic models ----------------------
class MatchModel(BaseModel): class MatchModel(BaseModel):
iif: Optional[str] iif: Optional[str]
@@ -56,11 +48,13 @@ class MatchModel(BaseModel):
tcp_dport: Optional[int] tcp_dport: Optional[int]
udp_dport: Optional[int] udp_dport: Optional[int]
class ActionModel(BaseModel): class ActionModel(BaseModel):
type: str # drop | accept | queue | redirect type: str # drop | accept | queue | redirect
queue_num: Optional[int] queue_num: Optional[int]
redirect_port: Optional[int] redirect_port: Optional[int]
class RuleModel(BaseModel): class RuleModel(BaseModel):
id: Optional[str] = Field(None, description="optional rule id; generated if missing") id: Optional[str] = Field(None, description="optional rule id; generated if missing")
family: Optional[str] = Field(DEFAULT_FAMILY) family: Optional[str] = Field(DEFAULT_FAMILY)
@@ -69,35 +63,48 @@ class RuleModel(BaseModel):
match: MatchModel match: MatchModel
action: ActionModel action: ActionModel
# ---------------------- Utilities ----------------------
def run_nft(args: List[str]) -> subprocess.CompletedProcess: class ReplaceResult(BaseModel):
version: str
applied: bool
rules_count: int
# ---------------------- Utilities ----------------------
def run_nft(args: List[str]) -> Tuple[str, str]:
"""
Run nft with given args.
Returns (stdout, stderr).
Raises RuntimeError on non-zero exit with stderr included.
"""
logger.debug("running nft: %s", " ".join(["nft"] + args))
try: try:
return subprocess.run(["nft"] + args, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True) 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: except subprocess.CalledProcessError as e:
# raise with stderr for easier debugging logger.error("nft failed: %s -- %s", " ".join(e.cmd), e.stderr.strip())
raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}") raise RuntimeError(f"nft failed: {' '.join(e.cmd)} -- {e.stderr.strip()}")
def ensure_table_chain(family: str, table: str, chain: str): def ensure_table_chain(family: str, table: str, chain: str) -> None:
# create table if missing (ignore error if exists) # create table if missing (ignore error if exists)
try: try:
logger.info("ensuring table %s.%s exists", family, table)
run_nft(["add", "table", family, table]) run_nft(["add", "table", family, table])
except RuntimeError: except RuntimeError:
# already exists or failed; ignore existence error logger.debug("table %s.%s may already exist", family, table)
pass
# create chain if missing: forward chain with hook forward # create chain if missing: forward chain with hook forward
try: try:
# Note: use type filter hook forward priority 0 logger.info("ensuring chain %s in table %s exists", chain, table)
run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"]) run_nft(["add", "chain", family, table, chain, "{", "type", "filter", "hook", "forward", "priority", "0", ";", "}"])
except RuntimeError: except RuntimeError:
# ignore if exists logger.debug("chain %s in table %s may already exist", chain, table)
pass
def encode_rule_comment(rule: Dict[str, Any]) -> str: def encode_rule_comment(rule: Dict[str, Any]) -> str:
# store the rule JSON as base64 to avoid quoting/escaping issues inside nft comment # store rule JSON as base64 to avoid quoting/escaping issues inside nft comment
j = json.dumps(rule, separators=(",", ":")) j = json.dumps(rule, separators=(",", ":"))
b = base64.b64encode(j.encode()).decode() b = base64.b64encode(j.encode()).decode()
rid = rule.get("id") or "" rid = rule.get("id") or ""
@@ -115,28 +122,28 @@ def decode_comment_payload(comment: str) -> Optional[Dict[str, Any]]:
j = base64.b64decode(b64.encode()).decode() j = base64.b64decode(b64.encode()).decode()
return json.loads(j) return json.loads(j)
except Exception: except Exception:
logger.debug("failed to decode comment payload: %s", comment)
return None return None
def build_nft_match_expr(match: MatchModel) -> List[str]: def build_nft_match_fragment(match: MatchModel) -> List[str]:
expr: List[str] = [] frag: List[str] = []
if match.iif: if match.iif:
expr += ["iif", match.iif] frag += ["iif", f'"{match.iif}"']
if match.oif: if match.oif:
expr += ["oif", match.oif] frag += ["oif", f'"{match.oif}"']
if match.meta_length is not None: if match.meta_length is not None:
# accept either range string or int frag += ["meta", "length", str(match.meta_length)]
expr += ["meta", "length", str(match.meta_length)]
if match.ip_proto: if match.ip_proto:
expr += ["ip", "protocol", match.ip_proto] frag += ["ip", "protocol", match.ip_proto]
if match.tcp_dport: if match.tcp_dport:
expr += ["tcp", "dport", str(match.tcp_dport)] frag += ["tcp", "dport", str(match.tcp_dport)]
if match.udp_dport: if match.udp_dport:
expr += ["udp", "dport", str(match.udp_dport)] frag += ["udp", "dport", str(match.udp_dport)]
return expr return frag
def build_nft_action_expr(action: ActionModel) -> List[str]: def build_nft_action_fragment(action: ActionModel) -> List[str]:
if action.type == "drop": if action.type == "drop":
return ["drop"] return ["drop"]
if action.type == "accept": if action.type == "accept":
@@ -151,81 +158,77 @@ def build_nft_action_expr(action: ActionModel) -> List[str]:
raise ValueError(f"unsupported action type: {action.type}") raise ValueError(f"unsupported action type: {action.type}")
def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]: def nft_rule_line_from_model(rule: RuleModel) -> str:
# ensure table/chain """
ensure_table_chain(rule.family, rule.table, rule.chain) Produce a single-line nft command:
add rule <family> <table> <chain> <match...> <action...> comment "<encoded>"
# ensure id """
if not rule.id: match_frag = build_nft_match_fragment(rule.match)
rule.id = str(uuid.uuid4()) action_frag = build_nft_action_fragment(rule.action)
# build nft command args
match_expr = build_nft_match_expr(rule.match)
action_expr = build_nft_action_expr(rule.action)
comment = encode_rule_comment(rule.dict()) comment = encode_rule_comment(rule.dict())
args: List[str] = ["add", "rule", rule.family, rule.table, rule.chain] parts = ["add", "rule", rule.family, rule.table, rule.chain]
args += match_expr parts += match_frag
args += action_expr parts += action_frag
args += ["comment", comment] return " ".join(parts) + f' comment "{comment}"'
try:
run_nft(args)
return {"id": rule.id, "status": "added"}
except RuntimeError as e:
raise HTTPException(status_code=500, detail=str(e))
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]]:
try: try:
out = run_nft(["list", "table", family, table]).stdout out, _ = run_nft(["list", "table", family, table])
except RuntimeError: except RuntimeError:
logger.debug("no table %s.%s found when listing rules", family, table)
return [] return []
results: List[Dict[str, Any]] = [] results: List[Dict[str, Any]] = []
# naive parse: nft prints individual rules as lines; look for comment "mitm_json:"
for line in out.splitlines(): for line in out.splitlines():
line = line.strip() line = line.strip()
if "comment" in line and "mitm_json:" in line: if "comment" in line and "mitm_json:" in line:
# find comment payload part: comment "..."
# format often: "comment "mitm_id:... mitm_json:...""
try: try:
# extract between first pair of double quotes first_quote = line.index('"')
first_quote = line.index('\"') last_quote = line.rindex('"')
last_quote = line.rindex('\"')
comment_str = line[first_quote + 1:last_quote] comment_str = line[first_quote + 1:last_quote]
except ValueError: except ValueError:
# fallback: take substring after comment
comment_str = line.split("comment", 1)[1].strip() comment_str = line.split("comment", 1)[1].strip()
parsed = decode_comment_payload(comment_str) parsed = decode_comment_payload(comment_str)
if parsed is not None: if parsed is not None:
results.append(parsed) results.append(parsed)
logger.debug("listed %d nft rules from %s.%s", len(results), family, table)
return results return results
def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]: def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]:
"""Delete rules whose comment contains mitm_id in ids. ensure_table_chain(rule.family, rule.table, rule.chain)
Returns list of deleted ids. if not rule.id:
""" rule.id = str(uuid.uuid4())
cmd_text = nft_rule_line_from_model(rule)
try: try:
out = run_nft(["list", "chain", family, table, chain, "-a"]).stdout # 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: except RuntimeError:
logger.debug("no chain %s in table %s.%s when attempting delete", chain, family, table)
return [] return []
deleted: List[str] = [] deleted: List[str] = []
# nft -a prints rules; each rule line may contain a comment and a trailing handle number: "... comment \\"mitm_id:...\\" ... handle 5"
for line in out.splitlines(): for line in out.splitlines():
if "comment" in line and "mitm_id:" in line: if "comment" in line and "mitm_id:" in line:
# extract comment between quotes
try: try:
q1 = line.index('\"') q1 = line.index('"')
q2 = line.index('\"', q1 + 1) q2 = line.index('"', q1 + 1)
comment_str = line[q1 + 1:q2] comment_str = line[q1 + 1:q2]
except ValueError: except ValueError:
# fallback: substring
comment_str = line.split("comment", 1)[1] comment_str = line.split("comment", 1)[1]
parsed = decode_comment_payload(comment_str) parsed = decode_comment_payload(comment_str)
@@ -233,10 +236,8 @@ def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) ->
continue continue
rid = parsed.get("id") rid = parsed.get("id")
if rid in ids: if rid in ids:
# find handle number
m = re.search(r"handle\s+(\d+)", line) m = re.search(r"handle\s+(\d+)", line)
if not m: if not m:
# try to find handle on the next token(s) - naive fallback
parts = line.split() parts = line.split()
if "handle" in parts: if "handle" in parts:
hi = parts.index("handle") hi = parts.index("handle")
@@ -251,24 +252,23 @@ def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) ->
try: try:
run_nft(["delete", "rule", family, table, chain, "handle", handle]) run_nft(["delete", "rule", family, table, chain, "handle", handle])
logger.info("deleted nft rule id=%s handle=%s", rid, handle)
deleted.append(rid) deleted.append(rid)
except RuntimeError: except RuntimeError:
# ignore deletion errors for now logger.error("failed to delete nft rule id=%s handle=%s", rid, handle)
continue continue
return deleted return deleted
# ---------------------- API endpoints ---------------------- # ---------------------- API endpoints ----------------------
router = APIRouter()
@router.get("/rules") @router.get("/rules")
def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE): def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_TABLE):
rules = list_rules_from_nft(family, table) rules = list_rules_from_nft(family, table)
return {"count": len(rules), "rules": rules} return {"count": len(rules), "rules": rules, "version": _current_version}
@router.post("/rules") @router.post("/rules")
def post_rules(payload: Union[RuleModel, List[RuleModel]]): def post_rules(payload: Union[RuleModel, List[RuleModel]]):
# accept either single or list
rules = payload if isinstance(payload, list) else [payload] rules = payload if isinstance(payload, list) else [payload]
results = [] results = []
for r in rules: for r in rules:
@@ -276,12 +276,88 @@ def post_rules(payload: Union[RuleModel, List[RuleModel]]):
results.append(res) results.append(res)
return {"results": results} return {"results": results}
@router.delete("/rules") @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): def delete_rules(payload: Union[str, List[str]], family: Optional[str] = DEFAULT_FAMILY,
"""Delete one or many rules by id. table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
Payload can be a single id string or a list of ids.
"""
ids = [payload] if isinstance(payload, str) else payload ids = [payload] if isinstance(payload, str) else payload
deleted = delete_rules_by_ids(ids, family, table, chain) deleted = delete_rules_by_ids(ids, family, table, chain)
results = [{"id": i, "deleted": i in deleted} for i in ids] results = [{"id": i, "deleted": i in deleted} for i in ids]
return {"results": results} return {"results": results}
@router.put("/rules", response_model=ReplaceResult)
def put_rules(rules: List[RuleModel], family: Optional[str] = DEFAULT_FAMILY,
table: Optional[str] = DEFAULT_TABLE, chain: Optional[str] = DEFAULT_CHAIN):
"""
Replace entire ordered rule set by generating an nft script and applying via `nft -f`.
Behavior:
1. Validate rules and ensure family/table/chain match (if provided).
2. Ensure the table/chain exist.
3. Generate nft script that flushes the chain and adds rules in given order.
4. Write script to a secure temp file and run `nft -f <file>`.
5. On success update in-memory version token and return it.
Note: nft runs script sequentially; if nft errors mid-script, partial state may exist.
For stricter atomicity, implement the temp-table swap approach.
"""
# validate per-rule family/table/chain if present
for r in rules:
if r.family and r.family != family:
raise HTTPException(status_code=400, detail=f"rule family mismatch: {r.family} != {family}")
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
ensure_table_chain(family, table, chain)
# ensure rule IDs
for r in rules:
if not r.id:
r.id = str(uuid.uuid4())
# build nft script lines
lines: List[str] = []
# flush chain (clear existing ordered rules)
lines.append(f"flush chain {family} {table} {chain}")
# add ordered rules
for r in rules:
try:
lines.append(nft_rule_line_from_model(r))
except Exception as e:
raise HTTPException(status_code=400, detail=f"invalid rule: {e}")
script = "\n".join(lines) + "\n"
tmpfile_path: Optional[str] = None
try:
with tempfile.NamedTemporaryFile(mode="w", delete=False, prefix="mitm_nft_", suffix=".nft") as tf:
tmpfile_path = tf.name
tf.write(script)
tf.flush()
os.fsync(tf.fileno())
logger.info("wrote nft script to %s; applying...", tmpfile_path)
# apply the file
try:
run_nft(["-f", tmpfile_path])
except RuntimeError as e:
logger.error("failed applying nft script: %s", 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())
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)