test12345
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 9s

This commit is contained in:
2026-01-11 17:23:34 +01:00
parent 57de326910
commit f422a2979f

View File

@@ -1,58 +1,62 @@
# fastapi_nft_router.py
# -*- coding: utf-8 -*-
""" """
FastAPI router specialized for managing nftables rules for the FastAPI router for managing nftables rules for the 'bridge' family (bridge),
'bridge' family, 'filter' table, 'forward' chain. default table 'mitm_tbl' and default chain 'forward'.
This enhanced version uses Python enums and narrowly-typed Pydantic Features:
models wherever sensible so the frontend can request 'options' and - Typed Pydantic models and enums to help the frontend populate dropdowns.
present user-friendly dropdowns. It keeps the dynamic `expr` model but - /options endpoint to return enum choices for UI selects.
provides many structured expression types (meta, ether, ip, tcp/udp, - List / preview / add / delete / update endpoints for rules.
ct, verdict, raw) implemented as discriminated unions. - Improved parsing of `nft --json list ruleset` (returns position, handle,
comment, verdict, verdict_details, exprs, nft_rule).
- Automatic creation of missing table and chain with conservative defaults.
- Structured, informative logging suitable for debugging and audit.
Features added in this version: Security:
- Rich enums for fields (MetaKey, EtherField, IPDir, Ops, Verdict, Proto, - This router issues `nft` commands on the host. Run the FastAPI process as root
ConntrackState, RejectType, IcmpType, LogGroup). or grant the binary / process the appropriate capabilities (NET_ADMIN).
- FastAPI endpoint `/options` that returns allowed enum choices for the - Consider restricting access to these endpoints (authentication, network restrictions)
frontend to populate dropdowns. before exposing on a network.
- Strict Pydantic models with `kind` discriminators for expression types.
- Validation to ensure family==bridge and chain/table default to Mounting:
filter/forward but still configurable if needed. from fastapi import FastAPI
- Preview and apply endpoints unchanged in behavior but now accept from fastapi_nft_router import router as nft_router
enumerated inputs which make the generated nft syntax safer and
simpler to render in the UI. app = FastAPI()
- Improved `nft_list_rules` parsing: returns rule `handle`, `position` app.include_router(nft_router)
(1-based within chain), `comment`, and best-effort `verdict` so the
frontend can present accurate update/delete controls.
Security: still run as root or with nft capabilities.
""" """
import logging
from typing import Any, Dict, List, Optional, Union, Literal from typing import Any, Dict, List, Optional, Union, Literal
import subprocess import subprocess
import shutil import shutil
import json import json
import logging
from enum import Enum from enum import Enum
from fastapi import APIRouter, HTTPException, Body from fastapi import APIRouter, HTTPException, Body
from pydantic import BaseModel, Field, validator from pydantic import BaseModel, Field, validator
# Router + logger # Router & logging -------------------------------------------------------
router = APIRouter() router = APIRouter(prefix="/api/nft/bridge/forward", tags=["nftables"])
logger = logging.getLogger("nftables") logger = logging.getLogger("nftables")
# Default logger configuration is left to the main application; here we
# emit messages at DEBUG/INFO/ERROR as appropriate.
logger.debug("nftables router module loaded") logger.debug("nftables router module loaded")
# Defaults # Defaults & local nft binary detection ---------------------------------
DEFAULT_TABLE = "mitm_tbl" DEFAULT_TABLE = "mitm_tbl"
DEFAULT_CHAIN = "forward" DEFAULT_CHAIN = "forward"
DEFAULT_FAMILY = "bridge" DEFAULT_FAMILY = "bridge"
NFT_BIN = shutil.which("nft") NFT_BIN = shutil.which("nft")
# ------------------ Enums (for front-end dropdowns) ------------------ # ------------------ Enums (for front-end dropdowns & typing) ------------
class Family(str, Enum): class Family(str, Enum):
bridge = "bridge" bridge = DEFAULT_FAMILY
class Table(str, Enum): class Table(str, Enum):
filter = "filter" table = DEFAULT_TABLE
raw = "raw"
class Chain(str, Enum): class Chain(str, Enum):
forward = "forward" forward = "forward"
@@ -101,7 +105,6 @@ class ConntrackState(str, Enum):
class RejectType(str, Enum): class RejectType(str, Enum):
icmp = "icmp" icmp = "icmp"
tcp_reset = "tcp reset" tcp_reset = "tcp reset"
# user may also supply raw later
class IcmpType(str, Enum): class IcmpType(str, Enum):
dest_unreachable = "destination-unreachable" dest_unreachable = "destination-unreachable"
@@ -122,92 +125,237 @@ class LogGroup(int, Enum):
g6 = 6 g6 = 6
g7 = 7 g7 = 7
# ------------------ Pydantic expression models (discriminated unions) ------------------ # ------------------ Pydantic expression models (discriminated unions) ---
class BaseExpr(BaseModel): class BaseExpr(BaseModel):
"""Base expression with a discriminator 'kind'."""
kind: str kind: str
class Config: class Config:
extra = "forbid" extra = "forbid"
class MetaExpr(BaseExpr): class MetaExpr(BaseExpr):
kind: Literal["meta"] = Field(default="meta") kind: Literal["meta"] = Field(default="meta")
key: MetaKey key: MetaKey
op: Op = Op.eq op: Op = Op.eq
value: str value: str
class EtherExpr(BaseExpr): class EtherExpr(BaseExpr):
kind: Literal["ether"] = Field(default="ether") kind: Literal["ether"] = Field(default="ether")
field: EtherField field: EtherField
op: Op = Op.eq op: Op = Op.eq
value: str value: str
class IPExpr(BaseExpr): class IPExpr(BaseExpr):
kind: Literal["ip"] = Field(default="ip") kind: Literal["ip"] = Field(default="ip")
side: IPDir side: IPDir
op: Op = Op.eq op: Op = Op.eq
value: str value: str
class ProtoPortExpr(BaseExpr): class ProtoPortExpr(BaseExpr):
kind: Literal["l4"] = Field(default="l4") kind: Literal["l4"] = Field(default="l4")
proto: Proto proto: Proto
sport: Optional[str] = None sport: Optional[str] = None
dport: Optional[str] = None dport: Optional[str] = None
class CTEexpr(BaseExpr): class CTEexpr(BaseExpr):
kind: Literal["ct"] = Field(default="ct") kind: Literal["ct"] = Field(default="ct")
state: ConntrackState state: ConntrackState
class VerdictExpr(BaseExpr): class VerdictExpr(BaseExpr):
kind: Literal["verdict"] = Field(default="verdict") kind: Literal["verdict"] = Field(default="verdict")
verdict: Verdict verdict: Verdict
class RejectExpr(BaseExpr): class RejectExpr(BaseExpr):
kind: Literal["reject"] = Field(default="reject") kind: Literal["reject"] = Field(default="reject")
reject_type: RejectType reject_type: RejectType
icmp_type: Optional[IcmpType] = None icmp_type: Optional[IcmpType] = None
class LogExpr(BaseExpr): class LogExpr(BaseExpr):
kind: Literal["log"] = Field(default="log") kind: Literal["log"] = Field(default="log")
prefix: Optional[str] = None prefix: Optional[str] = None
group: Optional[LogGroup] = None group: Optional[LogGroup] = None
class RawExpr(BaseExpr): class RawExpr(BaseExpr):
kind: Literal["raw"] = Field(default="raw") kind: Literal["raw"] = Field(default="raw")
snippet: str snippet: str
# Combine into a discriminated union
Expr = Union[MetaExpr, EtherExpr, IPExpr, ProtoPortExpr, CTEexpr, VerdictExpr, RejectExpr, LogExpr, RawExpr]
# ------------------ Rule model ------------------ # Union of accepted expression types for request validation
Expr = Union[
MetaExpr,
EtherExpr,
IPExpr,
ProtoPortExpr,
CTEexpr,
VerdictExpr,
RejectExpr,
LogExpr,
RawExpr,
]
# ------------------ Rule model -----------------------------------------
class RuleModel(BaseModel): class RuleModel(BaseModel):
"""Model for creating/inserting rules through the API."""
family: Family = Family.bridge family: Family = Family.bridge
table: Table = Table.filter table: Table = Table.table
chain: Chain = Chain.forward chain: Chain = Chain.forward
expr: List[Expr] = Field(default_factory=list) expr: List[Expr] = Field(default_factory=list)
comment: Optional[str] = None comment: Optional[str] = None
position: Optional[int] = None # position to insert (1-based) position: Optional[int] = None # 1-based insert position (if provided)
handle: Optional[int] = None handle: Optional[int] = None
@validator("family") @validator("family")
def only_bridge(cls, v): def only_bridge(cls, v: Family) -> Family:
if v != Family.bridge: if v != Family.bridge:
raise ValueError("This router only manages family 'bridge'") raise ValueError("This router only manages family 'bridge'")
return v return v
# ------------------ Helpers ------------------
def ensure_nft_available(): # ------------------ Low-level helpers -----------------------------------
def ensure_nft_available() -> None:
"""Raise HTTPException if nft binary is not found on the host."""
if not NFT_BIN: if not NFT_BIN:
logger.error("nft binary not found on server")
raise HTTPException(status_code=500, detail="nft binary not found on server") raise HTTPException(status_code=500, detail="nft binary not found on server")
def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, Any]: def run_nft_cmd(cmd: str) -> Dict[str, Any]:
"""
Execute an nft script containing a single command `cmd` using `nft -f -`.
Returns a dict with stdout/stderr. Raises HTTPException on error.
Logging:
- Info-level for run attempts and success.
- Debug-level for full stdout/stderr content.
- Error-level on failure with nft stderr included.
"""
ensure_nft_available() ensure_nft_available()
cmd = [NFT_BIN, "--json", "list", "ruleset"] full_cmd = [NFT_BIN, "-f", "-"]
script = cmd.rstrip() + "\n" # ensure newline
logger.info("Running nft command: %s", cmd)
logger.debug("Executing: %s ; script: %s", full_cmd, script)
try: try:
out = subprocess.check_output(cmd, stderr=subprocess.PIPE) proc = subprocess.run(
full_cmd,
input=script.encode(),
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
check=True,
)
stdout = proc.stdout.decode()
stderr = proc.stderr.decode()
logger.info("nft command succeeded (stdout %d bytes, stderr %d bytes)", len(stdout), len(stderr))
logger.debug("nft stdout: %s", stdout or "<empty>")
if stderr:
logger.debug("nft stderr: %s", stderr)
return {"stdout": stdout, "stderr": stderr}
except subprocess.CalledProcessError as e:
err = e.stderr.decode() if e.stderr else str(e)
logger.error("nft command failed: %s", err)
raise HTTPException(status_code=500, detail=err)
def ensure_table_and_chain_exist(family: str, table: str, chain: str) -> None:
"""
Ensure the given table and chain exist for `family`. Create them if missing.
Creation defaults:
- add table <family> <table>
- add chain <family> <table> <chain> { type filter hook <chain> priority 0; policy accept; }
for chain in (input, forward, output). For other chain names a simple chain is created
without hook: `add chain <family> <table> <chain> { policy accept; }`.
Raises HTTPException on nft failure.
"""
logger.debug("Checking existence: family=%s table=%s chain=%s", family, table, chain)
ensure_nft_available()
try:
out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE)
parsed = json.loads(out) parsed = json.loads(out)
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
logger.error("Failed to list ruleset: %s", e.stderr.decode())
raise HTTPException(status_code=500, detail=f"nft failed: {e.stderr.decode()}")
items = parsed.get("nftables") if isinstance(parsed, dict) else parsed
if not isinstance(items, list):
items = []
table_exists = False
chain_exists = False
for item in items:
if "table" in item:
t = item["table"]
if isinstance(t, dict) and t.get("name") == table and t.get("family") == family:
table_exists = True
if "chain" in item:
ch = item["chain"]
if isinstance(ch, dict) and ch.get("name") == chain and ch.get("table") == table and ch.get("family") == family:
chain_exists = True
if not table_exists:
logger.info("Table '%s' (family=%s) not found. Creating.", table, family)
try:
cmd = f"add table {family} {table}"
logger.debug("Creating table with: %s", cmd)
run_nft_cmd(cmd)
logger.info("Table '%s' created.", table)
except HTTPException as e:
logger.error("Failed to create table '%s': %s", table, getattr(e, "detail", str(e)))
raise
if not chain_exists:
logger.info("Chain '%s' in table '%s' (family=%s) not found. Creating.", chain, table, family)
try:
if chain in ("input", "forward", "output"):
# create base chain with hook
cmd = (
f"add chain {family} {table} {chain} {{ "
f"type filter hook {chain} priority 0; policy accept; }}"
)
else:
# create a user chain (no hook)
cmd = f"add chain {family} {table} {chain} {{ policy accept; }}"
logger.debug("Creating chain with: %s", cmd)
run_nft_cmd(cmd)
logger.info("Chain '%s' created in table '%s'.", chain, table)
except HTTPException as e:
logger.error("Failed to create chain '%s' in table '%s': %s", chain, table, getattr(e, "detail", str(e)))
raise
def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, Any]:
"""
Parse `nft --json list ruleset` and return a list of rules with metadata.
Each returned rule includes:
- family, table, chain
- handle (if present)
- position (1-based within its chain)
- comment (if present, best-effort)
- verdict (best-effort string: 'accept'/'drop'/'reject' or None)
- verdict_details (raw nested object for reject or other complex verdicts)
- exprs (original expression list from nft JSON)
- nft_rule (the original rule dict from nft JSON)
This function is conservative and aims to give the UI enough info to
display and edit rules precisely (by handle or position).
"""
ensure_nft_available()
try:
out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE)
parsed = json.loads(out)
except subprocess.CalledProcessError as e:
logger.error("Failed to list ruleset: %s", e.stderr.decode())
raise HTTPException(status_code=500, detail=f"nft failed: {e.stderr.decode()}") raise HTTPException(status_code=500, detail=f"nft failed: {e.stderr.decode()}")
results: List[Dict[str, Any]] = [] results: List[Dict[str, Any]] = []
@@ -218,7 +366,6 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
items = [] items = []
for item in items: for item in items:
# skip table/chain metadata entries (we only need rule entries)
if "rule" not in item: if "rule" not in item:
continue continue
@@ -232,37 +379,31 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
position = counters[key] position = counters[key]
handle = r.get("handle") handle = r.get("handle")
exprs = r.get("expr", []) # original expression list from nft JSON exprs = r.get("expr", []) # original expression list
comment = None comment: Optional[str] = None
verdict = None verdict: Optional[str] = None
verdict_details = None verdict_details: Optional[Any] = None
# scan expressions to extract comment and verdict/action # scan expressions to extract comment and verdict/action
for expr in exprs: for expr in exprs:
if not isinstance(expr, dict): if not isinstance(expr, dict):
continue continue
# comment can appear as {"comment":"text"} or {"comment": {"text": "..."}} depending on nft json variant
if "comment" in expr: if "comment" in expr:
# handle both simple and nested forms
c = expr.get("comment") c = expr.get("comment")
if isinstance(c, str): if isinstance(c, str):
comment = c comment = c
elif isinstance(c, dict): elif isinstance(c, dict):
# some representations: {"comment": {"text": "..."}} or {"comment": {"str": "..."}}
comment = c.get("text") or c.get("str") or comment comment = c.get("text") or c.get("str") or comment
# verdict forms # common verdict shapes: {"verdict": {"accept": null}} or {"drop": null}
if "verdict" in expr: if "verdict" in expr:
v = expr["verdict"] v = expr["verdict"]
# v often looks like {"accept": None} or {"drop": None} or {"reject": {...}}
if isinstance(v, dict): if isinstance(v, dict):
# take the first key as the action
k = next(iter(v.keys()), None) k = next(iter(v.keys()), None)
verdict = k verdict = k
verdict_details = v.get(k) verdict_details = v.get(k)
else: else:
verdict = str(v) verdict = str(v)
# older/alternate forms
if "drop" in expr and verdict is None: if "drop" in expr and verdict is None:
verdict = "drop" verdict = "drop"
verdict_details = expr.get("drop") verdict_details = expr.get("drop")
@@ -273,18 +414,20 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
verdict = "reject" verdict = "reject"
verdict_details = expr.get("reject") verdict_details = expr.get("reject")
results.append({ results.append(
"family": family, {
"table": table_name, "family": family,
"chain": chain_name, "table": table_name,
"handle": handle, "chain": chain_name,
"position": position, "handle": handle,
"comment": comment, "position": position,
"verdict": verdict, # e.g. "accept", "drop", "reject", or None "comment": comment,
"verdict_details": verdict_details, # raw details for rejects/other actions "verdict": verdict,
"exprs": exprs, # original expression list (renamed, clearer) "verdict_details": verdict_details,
"nft_rule": r, # the original nft JSON dict for this rule "exprs": exprs,
}) "nft_rule": r,
}
)
# filter by requested table/chain if provided # filter by requested table/chain if provided
if table or chain: if table or chain:
@@ -292,9 +435,8 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
return {"rules": results} return {"rules": results}
def expr_to_nft_snippet(e: Expr) -> str: def expr_to_nft_snippet(e: Expr) -> str:
# runtime dispatch via model type """Convert a typed Expr into a short nft syntax snippet (used for preview/add)."""
if isinstance(e, MetaExpr): if isinstance(e, MetaExpr):
val = e.value val = e.value
key = e.key.value key = e.key.value
@@ -318,8 +460,7 @@ def expr_to_nft_snippet(e: Expr) -> str:
if e.reject_type == RejectType.icmp: if e.reject_type == RejectType.icmp:
t = f"reject with icmp type {e.icmp_type.value}" if e.icmp_type else "reject" t = f"reject with icmp type {e.icmp_type.value}" if e.icmp_type else "reject"
return t return t
else: return e.reject_type.value
return e.reject_type.value
if isinstance(e, LogExpr): if isinstance(e, LogExpr):
parts = ["log"] parts = ["log"]
if e.prefix: if e.prefix:
@@ -333,6 +474,7 @@ def expr_to_nft_snippet(e: Expr) -> str:
def rule_to_nft_cmd(rule: RuleModel) -> str: def rule_to_nft_cmd(rule: RuleModel) -> str:
"""Build the nft command string for add/insert operations from a RuleModel."""
expr_snippets = [expr_to_nft_snippet(e) for e in rule.expr] expr_snippets = [expr_to_nft_snippet(e) for e in rule.expr]
body = " ".join(s for s in expr_snippets if s) body = " ".join(s for s in expr_snippets if s)
if rule.position is not None: if rule.position is not None:
@@ -344,21 +486,13 @@ def rule_to_nft_cmd(rule: RuleModel) -> str:
return cmd return cmd
def run_nft_cmd(cmd: str) -> Dict[str, Any]: # ------------------ Endpoints -------------------------------------------
ensure_nft_available()
full_cmd = [NFT_BIN, "-f", "-"]
script = cmd + ""
try:
proc = subprocess.run(full_cmd, input=script.encode(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=True)
return {"stdout": proc.stdout.decode(), "stderr": proc.stderr.decode()}
except subprocess.CalledProcessError as e:
raise HTTPException(status_code=500, detail=(e.stderr.decode() or str(e)))
# ------------------ Endpoints ------------------
@router.get("/options") @router.get("/options")
def get_options(): def get_options() -> Dict[str, Any]:
"""Return available enum choices for the frontend dropdowns.""" """
Return allowed enum choices for the frontend dropdowns.
The frontend should call this once and cache results to populate selects.
"""
return { return {
"family": [f.value for f in Family], "family": [f.value for f in Family],
"table": [t.value for t in Table], "table": [t.value for t in Table],
@@ -375,42 +509,79 @@ def get_options():
"log_groups": [int(g.value) for g in LogGroup], "log_groups": [int(g.value) for g in LogGroup],
} }
@router.get("/rules") @router.get("/rules")
def list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN): def list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
"""
List rules for table+chain. Ensures the table/chain exist (creates if missing).
Returns a dict { "rules": [ { family, table, chain, handle, position, comment, verdict, verdict_details, exprs, nft_rule }, ... ] }
"""
# Ensure the namespace exists (create table/chain if missing)
ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain)
return nft_list_rules(table=table, chain=chain) return nft_list_rules(table=table, chain=chain)
@router.post("/rules/preview") @router.post("/rules/preview")
def preview_rule(rule: RuleModel = Body(...)): def preview_rule(rule: RuleModel = Body(...)) -> Dict[str, str]:
"""
Return the nft command that would be executed for the provided rule (without applying it).
Useful for the frontend preview step.
"""
try: try:
cmd = rule_to_nft_cmd(rule) cmd = rule_to_nft_cmd(rule)
except Exception as e: except Exception as e:
logger.error("Preview building failed: %s", str(e))
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
return {"cmd": cmd} return {"cmd": cmd}
@router.post("/rules") @router.post("/rules")
def add_rule(rule: RuleModel = Body(...)): def add_rule(rule: RuleModel = Body(...)) -> Dict[str, Any]:
"""
Insert (position) or append a rule into the specified table/chain.
Will create table/chain if missing.
"""
ensure_table_and_chain_exist(rule.family.value, rule.table.value, rule.chain.value)
cmd = rule_to_nft_cmd(rule) cmd = rule_to_nft_cmd(rule)
return run_nft_cmd(cmd) return run_nft_cmd(cmd)
@router.delete("/rules/{handle}") @router.delete("/rules/{handle}")
def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN): def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
ensure_nft_available() """
Delete a rule by nft handle. Ensures table/chain exist before attempting delete.
"""
ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain)
cmd = f"delete rule {table} {chain} handle {handle}" cmd = f"delete rule {table} {chain} handle {handle}"
return run_nft_cmd(cmd) return run_nft_cmd(cmd)
@router.put("/rules/{handle}") @router.put("/rules/{handle}")
def update_rule(handle: int, rule: RuleModel = Body(...), table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN): def update_rule(
# Delete by handle then insert at the same position (if available via list parse) handle: int,
# We attempt to find the position of the handle so the replacement keeps position. rule: RuleModel = Body(...),
table: str = DEFAULT_TABLE,
chain: str = DEFAULT_CHAIN,
) -> Dict[str, Any]:
"""
Replace a rule identified by its nft handle:
- Attempts to find the position of the given handle and insert the replacement at the same position.
- If position cannot be determined will append the replacement.
"""
ensure_table_and_chain_exist(rule.family.value, rule.table.value, rule.chain.value)
# Attempt to find existing position of handle
rules_info = nft_list_rules(table=table, chain=chain) rules_info = nft_list_rules(table=table, chain=chain)
position = None position: Optional[int] = None
for r in rules_info.get('rules', []): for r in rules_info.get("rules", []):
if r.get('handle') == handle: if r.get("handle") == handle:
position = r.get('position') position = r.get("position")
break break
# delete by handle
# delete existing rule by handle
delete_rule(handle, table=table, chain=chain) delete_rule(handle, table=table, chain=chain)
# if we found position, insert at that position; otherwise append
# insert replacement at found position if available
if position is not None: if position is not None:
rule.position = position rule.position = position
return add_rule(rule) return add_rule(rule)