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.
- 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
nftables router for FastAPI to manage nftables bridge rules dynamically.
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:
- Runs nft(8) commands; the process must have sufficient privileges (run as root or via sudo).
- This implementation stores the full rule JSON inside the nft rule comment as base64 to allow round-trip parsing.
- The API keeps sniffer separate; this module only manages kernel 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)
- Must run with privileges to call `nft` (root or via sudo).
- Rules are stored in the nft rule comment as base64-encoded JSON for round-trip parsing.
- PUT /rules writes a temporary nft script and runs `nft -f <file>`; this applies the new ordered rules.
"""
from fastapi import APIRouter, APIRouter, HTTPException
from fastapi import APIRouter, HTTPException
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 uuid
import json
import base64
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_CHAIN = "forward"
DEFAULT_FAMILY = "bridge"
# in-memory version token updated on successful PUT
_current_version: Optional[str] = None
# ---------------------- Pydantic models ----------------------
class MatchModel(BaseModel):
iif: Optional[str]
@@ -56,11 +48,13 @@ class MatchModel(BaseModel):
tcp_dport: Optional[int]
udp_dport: Optional[int]
class ActionModel(BaseModel):
type: str # drop | accept | queue | redirect
queue_num: Optional[int]
redirect_port: Optional[int]
class RuleModel(BaseModel):
id: Optional[str] = Field(None, description="optional rule id; generated if missing")
family: Optional[str] = Field(DEFAULT_FAMILY)
@@ -69,35 +63,48 @@ class RuleModel(BaseModel):
match: MatchModel
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:
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:
# 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()}")
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)
try:
logger.info("ensuring table %s.%s exists", family, table)
run_nft(["add", "table", family, table])
except RuntimeError:
# already exists or failed; ignore existence error
pass
logger.debug("table %s.%s may already exist", family, table)
# create chain if missing: forward chain with hook forward
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", ";", "}"])
except RuntimeError:
# ignore if exists
pass
logger.debug("chain %s in table %s may already exist", chain, table)
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=(",", ":"))
b = base64.b64encode(j.encode()).decode()
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()
return json.loads(j)
except Exception:
logger.debug("failed to decode comment payload: %s", comment)
return None
def build_nft_match_expr(match: MatchModel) -> List[str]:
expr: List[str] = []
def build_nft_match_fragment(match: MatchModel) -> List[str]:
frag: List[str] = []
if match.iif:
expr += ["iif", match.iif]
frag += ["iif", f'"{match.iif}"']
if match.oif:
expr += ["oif", match.oif]
frag += ["oif", f'"{match.oif}"']
if match.meta_length is not None:
# accept either range string or int
expr += ["meta", "length", str(match.meta_length)]
frag += ["meta", "length", str(match.meta_length)]
if match.ip_proto:
expr += ["ip", "protocol", match.ip_proto]
frag += ["ip", "protocol", match.ip_proto]
if match.tcp_dport:
expr += ["tcp", "dport", str(match.tcp_dport)]
frag += ["tcp", "dport", str(match.tcp_dport)]
if match.udp_dport:
expr += ["udp", "dport", str(match.udp_dport)]
return expr
frag += ["udp", "dport", str(match.udp_dport)]
return frag
def build_nft_action_expr(action: ActionModel) -> List[str]:
def build_nft_action_fragment(action: ActionModel) -> List[str]:
if action.type == "drop":
return ["drop"]
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}")
def add_rule_to_nft(rule: RuleModel) -> Dict[str, Any]:
# ensure table/chain
ensure_table_chain(rule.family, rule.table, rule.chain)
# ensure id
if not rule.id:
rule.id = str(uuid.uuid4())
# build nft command args
match_expr = build_nft_match_expr(rule.match)
action_expr = build_nft_action_expr(rule.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())
args: List[str] = ["add", "rule", rule.family, rule.table, rule.chain]
args += match_expr
args += action_expr
args += ["comment", comment]
try:
run_nft(args)
return {"id": rule.id, "status": "added"}
except RuntimeError as e:
raise HTTPException(status_code=500, detail=str(e))
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]]:
try:
out = run_nft(["list", "table", family, table]).stdout
out, _ = run_nft(["list", "table", family, table])
except RuntimeError:
logger.debug("no table %s.%s found when listing rules", family, table)
return []
results: List[Dict[str, Any]] = []
# naive parse: nft prints individual rules as lines; look for comment "mitm_json:"
for line in out.splitlines():
line = line.strip()
if "comment" in line and "mitm_json:" in line:
# find comment payload part: comment "..."
# format often: "comment "mitm_id:... mitm_json:...""
try:
# extract between first pair of double quotes
first_quote = line.index('\"')
last_quote = line.rindex('\"')
first_quote = line.index('"')
last_quote = line.rindex('"')
comment_str = line[first_quote + 1:last_quote]
except ValueError:
# fallback: take substring after comment
comment_str = line.split("comment", 1)[1].strip()
parsed = decode_comment_payload(comment_str)
if parsed is not None:
results.append(parsed)
logger.debug("listed %d nft rules from %s.%s", len(results), family, table)
return results
def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) -> List[str]:
"""Delete rules whose comment contains mitm_id in ids.
Returns list of deleted ids.
"""
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:
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:
logger.debug("no chain %s in table %s.%s when attempting delete", chain, family, table)
return []
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():
if "comment" in line and "mitm_id:" in line:
# extract comment between quotes
try:
q1 = line.index('\"')
q2 = line.index('\"', q1 + 1)
q1 = line.index('"')
q2 = line.index('"', q1 + 1)
comment_str = line[q1 + 1:q2]
except ValueError:
# fallback: substring
comment_str = line.split("comment", 1)[1]
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
rid = parsed.get("id")
if rid in ids:
# find handle number
m = re.search(r"handle\s+(\d+)", line)
if not m:
# try to find handle on the next token(s) - naive fallback
parts = line.split()
if "handle" in parts:
hi = parts.index("handle")
@@ -251,24 +252,23 @@ def delete_rules_by_ids(ids: List[str], family: str, table: str, chain: str) ->
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:
# ignore deletion errors for now
logger.error("failed to delete nft rule id=%s handle=%s", rid, handle)
continue
return deleted
# ---------------------- API endpoints ----------------------
router = APIRouter()
@router.get("/rules")
def get_rules(family: Optional[str] = DEFAULT_FAMILY, table: Optional[str] = DEFAULT_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")
def post_rules(payload: Union[RuleModel, List[RuleModel]]):
# accept either single or list
rules = payload if isinstance(payload, list) else [payload]
results = []
for r in rules:
@@ -276,12 +276,88 @@ def post_rules(payload: Union[RuleModel, List[RuleModel]]):
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):
"""Delete one or many rules by id.
Payload can be a single id string or a list of ids.
"""
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)
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)