add logging and improve api for nftables
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
This commit is contained in:
@@ -1,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)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user