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

This commit is contained in:
2026-01-11 17:43:59 +01:00
parent 6347bcafdf
commit f0790904df

View File

@@ -1,57 +1,40 @@
# fastapi_nft_router.py # fastapi_nft_router.py
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
""" """
FastAPI router for managing nftables rules for the 'bridge' family (bridge), FastAPI router for nftables (bridge family). Returns both JSON and
default table 'mitm_tbl' and default chain 'forward'. human-readable textual representations for rules so the frontend can
display exact nft textual lines and programmatically safe metadata.
Features: Features:
- Typed Pydantic models and enums to help the frontend populate dropdowns. - nft_list_rules returns nft_rule_text_full, nft_rule_text (no handle), and add_command per rule.
- /options endpoint to return enum choices for UI selects. - Robust mapping by handle using both JSON and textual chain dump.
- List / preview / add / delete / update endpoints for rules. - Fallback JSON->text renderer for cases where textual mapping is missing.
- Improved parsing of `nft --json list ruleset` (returns position, handle, - Automatic create table/chain if missing, logging, and typed models for frontend.
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.
Security:
- This router issues `nft` commands on the host. Run the FastAPI process as root
or grant the binary / process the appropriate capabilities (NET_ADMIN).
- Consider restricting access to these endpoints (authentication, network restrictions)
before exposing on a network.
Mounting:
from fastapi import FastAPI
from fastapi_nft_router import router as nft_router
app = FastAPI()
app.include_router(nft_router)
""" """
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 import logging
import re
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 & logging ------------------------------------------------------- # Router & logging
router = APIRouter() router = APIRouter()
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 & local nft binary detection --------------------------------- # Defaults & nft binary locator
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 & typing) ------------ # ---------- Enums (for frontend dropdowns) ------------------------------
class Family(str, Enum): class Family(str, Enum):
bridge = DEFAULT_FAMILY bridge = DEFAULT_FAMILY
@@ -125,93 +108,71 @@ class LogGroup(int, Enum):
g6 = 6 g6 = 6
g7 = 7 g7 = 7
# ------------------ Pydantic expression models (discriminated unions) --- # ---------- Pydantic expression models ---------------------------------
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
# Union of accepted expression types for request validation
Expr = Union[ Expr = Union[
MetaExpr, MetaExpr, EtherExpr, IPExpr, ProtoPortExpr, CTEexpr,
EtherExpr, VerdictExpr, RejectExpr, LogExpr, RawExpr,
IPExpr,
ProtoPortExpr,
CTEexpr,
VerdictExpr,
RejectExpr,
LogExpr,
RawExpr,
] ]
# ---------- Rule model -------------------------------------------------
# ------------------ 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.table 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 # 1-based insert position (if provided) position: Optional[int] = None
handle: Optional[int] = None handle: Optional[int] = None
@validator("family") @validator("family")
@@ -220,41 +181,24 @@ class RuleModel(BaseModel):
raise ValueError("This router only manages family 'bridge'") raise ValueError("This router only manages family 'bridge'")
return v return v
# ---------- Low-level helpers -----------------------------------------
# ------------------ Low-level helpers -----------------------------------
def ensure_nft_available() -> None: 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") logger.error("nft binary not found")
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 run_nft_cmd(cmd: str) -> Dict[str, Any]: def run_nft_cmd(cmd: str) -> Dict[str, Any]:
""" """Run `nft -f -` with the given single-line script (ensures newline)."""
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()
full_cmd = [NFT_BIN, "-f", "-"] full_cmd = [NFT_BIN, "-f", "-"]
script = cmd.rstrip() + "\n" # ensure newline script = cmd.rstrip() + "\n"
logger.info("Running nft command: %s", cmd) logger.info("Running nft command: %s", cmd)
logger.debug("Executing: %s ; script: %s", full_cmd, script) logger.debug("Executing: %s ; script: %s", full_cmd, script)
try: try:
proc = subprocess.run( proc = subprocess.run(full_cmd, input=script.encode(), stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=True)
full_cmd,
input=script.encode(),
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
check=True,
)
stdout = proc.stdout.decode() stdout = proc.stdout.decode()
stderr = proc.stderr.decode() stderr = proc.stderr.decode()
logger.info("nft command succeeded (stdout %d bytes, stderr %d bytes)", len(stdout), len(stderr)) logger.info("nft command succeeded (%d bytes stdout, %d bytes stderr)", len(stdout), len(stderr))
logger.debug("nft stdout: %s", stdout or "<empty>") logger.debug("nft stdout: %s", stdout or "<empty>")
if stderr: if stderr:
logger.debug("nft stderr: %s", stderr) logger.debug("nft stderr: %s", stderr)
@@ -264,22 +208,10 @@ def run_nft_cmd(cmd: str) -> Dict[str, Any]:
logger.error("nft command failed: %s", err) logger.error("nft command failed: %s", err)
raise HTTPException(status_code=500, detail=err) raise HTTPException(status_code=500, detail=err)
def ensure_table_and_chain_exist(family: str, table: str, chain: str) -> None: def ensure_table_and_chain_exist(family: str, table: str, chain: str) -> None:
""" """Detect and create table/chain if missing (conservative defaults)."""
Ensure the given table and chain exist for `family`. Create them if missing. logger.debug("Checking/existence for family=%s table=%s chain=%s", family, table, chain)
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() ensure_nft_available()
try: try:
out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE) out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE)
parsed = json.loads(out) parsed = json.loads(out)
@@ -304,107 +236,158 @@ def ensure_table_and_chain_exist(family: str, table: str, chain: str) -> None:
chain_exists = True chain_exists = True
if not table_exists: if not table_exists:
logger.info("Table '%s' (family=%s) not found. Creating.", table, family) logger.info("Creating table %s %s", family, table)
try: run_nft_cmd(f"add table {family} {table}")
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: if not chain_exists:
logger.info("Chain '%s' in table '%s' (family=%s) not found. Creating.", chain, table, family) logger.info("Creating chain %s in table %s", chain, table)
try:
if chain in ("input", "forward", "output"): if chain in ("input", "forward", "output"):
# create base chain with hook run_nft_cmd(f"add chain {family} {table} {chain} {{ type filter hook {chain} priority 0; policy accept; }}")
cmd = (
f"add chain {family} {table} {chain} {{ "
f"type filter hook {chain} priority 0; policy accept; }}"
)
else: else:
# create a user chain (no hook) run_nft_cmd(f"add chain {family} {table} {chain} {{ policy accept; }}")
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_chain_text( # ---------- textual chain parser & JSON->text fallback -----------------
family: str, HANDLE_RE = re.compile(r"\bhandle\s+(\d+)\b", flags=re.IGNORECASE)
table: str,
chain: str, def nft_list_chain_text(family: str, table: str, chain: str) -> Dict[int, str]:
) -> Dict[int, str]:
""" """
Return a mapping: handle -> textual nft rule line Return mapping handle -> full textual line from:
extracted from `nft list chain <family> <table> <chain>`. nft list chain <family> <table> <chain>
Example return:
{
12: 'meta iifname "eth0" accept comment "allow lan"',
13: 'ip saddr 10.0.0.0/24 drop'
}
""" """
ensure_nft_available() ensure_nft_available()
cmd = [NFT_BIN, "list", "chain", family, table, chain] cmd = [NFT_BIN, "list", "chain", family, table, chain]
logger.debug("Listing chain in text mode: %s", " ".join(cmd)) logger.debug("Running textual chain list: %s", " ".join(cmd))
try: try:
out = subprocess.check_output(cmd, stderr=subprocess.PIPE).decode() out = subprocess.check_output(cmd, stderr=subprocess.PIPE).decode()
except subprocess.CalledProcessError as e: except subprocess.CalledProcessError as e:
logger.error("Failed to list chain text: %s", e.stderr.decode()) logger.error("Failed textual chain list: %s", e.stderr.decode())
# bubble up - caller may fallback; raising is acceptable here
raise HTTPException(status_code=500, detail=e.stderr.decode()) raise HTTPException(status_code=500, detail=e.stderr.decode())
rules: Dict[int, str] = {} mapping: Dict[int, str] = {}
for line in out.splitlines(): for line in out.splitlines():
line = line.strip() s = line.strip()
# Typical rule line contains: "handle <n>" if not s:
# Example: continue
# meta iifname "eth0" accept comment "foo" handle 7 m = HANDLE_RE.search(s)
if " handle " not in line: if not m:
continue continue
try: try:
rule_part, handle_part = line.rsplit(" handle ", 1) handle = int(m.group(1))
handle = int(handle_part.strip()) # full line including handle token
rules[handle] = rule_part.strip() mapping[handle] = s
except Exception: logger.debug("Text mapping: handle=%d -> %s", handle, s)
logger.debug("Could not parse rule line: %s", line) except Exception as ex:
logger.debug("Failed parsing handle from line: %s (%s)", s, ex)
continue
return mapping
return rules def json_exprs_to_text(exprs: List[Any]) -> str:
def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, Any]:
""" """
Parse nft rules using JSON mode for structure AND text mode for readability. Build a best-effort single-line textual rule clause from nft JSON exprs.
Used as a fallback when textual chain dump doesn't include the handle mapping.
"""
parts: List[str] = []
for ex in exprs:
if not isinstance(ex, dict):
continue
# comment
if "comment" in ex:
c = ex["comment"]
if isinstance(c, str):
parts.append(f'comment "{c}"')
elif isinstance(c, dict):
txt = c.get("text") or c.get("str")
if txt:
parts.append(f'comment "{txt}"')
continue
# verdict
if "verdict" in ex:
v = ex["verdict"]
if isinstance(v, dict):
k = next(iter(v.keys()), None)
parts.append(k if k else "verdict")
else:
parts.append(str(v))
continue
if "drop" in ex:
parts.append("drop")
continue
if "accept" in ex:
parts.append("accept")
continue
# match shapes
if "match" in ex:
m = ex["match"]
left = m.get("left")
op = m.get("op")
right = m.get("right")
if left and op and (right is not None):
parts.append(f"{left} {op} {right}")
continue
if "meta" in ex:
meta = ex["meta"]
key = meta.get("key") or meta.get("name")
op = meta.get("op", "==")
val = meta.get("value")
if key and val is not None:
parts.append(f"meta {key} {op} {val}")
continue
if "ct" in ex:
ct = ex["ct"]
if isinstance(ct, dict):
for k, v in ct.items():
parts.append(f"ct {k} {v}")
continue
if "payload" in ex:
parts.append(json.dumps(ex["payload"]))
continue
if "log" in ex:
lg = ex["log"]
piece = "log"
if isinstance(lg, dict):
if lg.get("prefix"):
piece += f' prefix "{lg.get("prefix")}"'
if lg.get("group") is not None:
piece += f' group {lg.get("group")}'
parts.append(piece)
continue
# fallback: compact JSON
parts.append(json.dumps(ex))
return " ".join(parts)
Returned fields per rule: # ---------- JSON rules parsing with textual injection ------------------
def nft_list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
"""
Return structured rules with both JSON exprs and textual representations.
Each rule entry contains:
- family, table, chain - family, table, chain
- handle - handle
- position (1-based) - position (1-based)
- comment (best-effort) - comment (best-effort)
- verdict (best-effort) - verdict (best-effort)
- verdict_details - verdict_details
- exprs (JSON expressions) - exprs (JSON)
- nft_rule (TEXTUAL rule line, exactly as nft prints it) - nft_rule_text_full (exact textual line from `nft list chain ...`, includes 'handle N' if present)
- nft_rule_text (textual clause without trailing 'handle N')
- add_command (best-effort `add rule <table> <chain> ...` command)
""" """
ensure_nft_available() ensure_nft_available()
try: try:
out = subprocess.check_output( out = subprocess.check_output([NFT_BIN, "--json", "list", "ruleset"], stderr=subprocess.PIPE)
[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 (json): %s", e.stderr.decode()) logger.error("Failed to list ruleset (json): %s", e.stderr.decode())
raise HTTPException(status_code=500, detail=e.stderr.decode()) raise HTTPException(status_code=500, detail=e.stderr.decode())
# 2) Text rules (human-readable) # get textual mapping for the chain (may raise; catch and fallback)
text_rules = nft_list_chain_text(DEFAULT_FAMILY, table, chain) text_map: Dict[int, str] = {}
try:
text_map = nft_list_chain_text(DEFAULT_FAMILY, table, chain)
logger.debug("Obtained textual mapping with %d entries", len(text_map))
except HTTPException as e:
logger.debug("Unable to get textual chain dump: %s; will fallback per-rule", getattr(e, "detail", str(e)))
results: List[Dict[str, Any]] = [] results: List[Dict[str, Any]] = []
counters: Dict[str, int] = {} counters: Dict[str, int] = {}
@@ -416,7 +399,6 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
for item in items: for item in items:
if "rule" not in item: if "rule" not in item:
continue continue
r = item["rule"] r = item["rule"]
family = r.get("family") family = r.get("family")
table_name = r.get("table") table_name = r.get("table")
@@ -440,14 +422,12 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
for expr in exprs: for expr in exprs:
if not isinstance(expr, dict): if not isinstance(expr, dict):
continue continue
if "comment" in expr: if "comment" in expr:
c = expr["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):
comment = c.get("text") or c.get("str") comment = c.get("text") or c.get("str")
if "verdict" in expr: if "verdict" in expr:
v = expr["verdict"] v = expr["verdict"]
if isinstance(v, dict): if isinstance(v, dict):
@@ -455,7 +435,6 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
verdict_details = v.get(verdict) verdict_details = v.get(verdict)
else: else:
verdict = str(v) verdict = str(v)
if "drop" in expr and verdict is None: if "drop" in expr and verdict is None:
verdict = "drop" verdict = "drop"
if "accept" in expr and verdict is None: if "accept" in expr and verdict is None:
@@ -463,8 +442,29 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
if "reject" in expr and verdict is None: if "reject" in expr and verdict is None:
verdict = "reject" verdict = "reject"
results.append( # textual resolution: prefer exact mapping by handle
{ full_text: Optional[str] = None
if handle is not None:
full_text = text_map.get(handle)
if full_text:
# derive no-handle version by stripping final ' handle N' if present
m = HANDLE_RE.search(full_text)
if m:
no_handle_text = full_text[: m.start()].strip()
else:
no_handle_text = full_text
logger.debug("Mapped handle %s -> textual full='%s'", handle, full_text)
else:
# fallback: build readable text from exprs
no_handle_text = json_exprs_to_text(exprs)
full_text = (no_handle_text + f" handle {handle}") if handle is not None else no_handle_text
logger.debug("Fallback textual for handle %s: %s", handle, no_handle_text)
# build an add_command (best-effort) - uses add rule <table> <chain> <clause>
add_cmd_clause = no_handle_text or ""
add_command = f"add rule {table_name} {chain_name} {add_cmd_clause}".strip()
results.append({
"family": family, "family": family,
"table": table_name, "table": table_name,
"chain": chain_name, "chain": chain_name,
@@ -474,17 +474,15 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A
"verdict": verdict, "verdict": verdict,
"verdict_details": verdict_details, "verdict_details": verdict_details,
"exprs": exprs, "exprs": exprs,
# 👇 THIS IS THE IMPORTANT CHANGE "nft_rule_text_full": full_text, # e.g. 'meta iifname "eth0" accept comment "x" handle 7'
"nft_rule": text_rules.get(handle), "nft_rule_text": no_handle_text, # e.g. 'meta iifname "eth0" accept comment "x"'
} "add_command": add_command, # e.g. 'add rule mitm_tbl forward meta iifname "eth0" accept ...'
) })
return {"rules": results} return {"rules": results}
# ---------- Expr -> nft snippet used for preview/add --------------------
def expr_to_nft_snippet(e: Expr) -> str: def expr_to_nft_snippet(e: Expr) -> str:
"""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
@@ -506,13 +504,12 @@ def expr_to_nft_snippet(e: Expr) -> str:
return e.verdict.value if e.verdict != Verdict.continue_ else "continue" return e.verdict.value if e.verdict != Verdict.continue_ else "continue"
if isinstance(e, RejectExpr): if isinstance(e, RejectExpr):
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" return f"reject with icmp type {e.icmp_type.value}" if e.icmp_type else "reject"
return t
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:
parts.append(f"prefix \"{e.prefix}\"") parts.append(f'prefix "{e.prefix}"')
if e.group is not None: if e.group is not None:
parts.append(f"group {int(e.group)}") parts.append(f"group {int(e.group)}")
return " ".join(parts) return " ".join(parts)
@@ -520,9 +517,7 @@ def expr_to_nft_snippet(e: Expr) -> str:
return e.snippet return e.snippet
raise ValueError("Unsupported expression type") raise ValueError("Unsupported expression type")
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:
@@ -530,17 +525,13 @@ def rule_to_nft_cmd(rule: RuleModel) -> str:
else: else:
cmd = f"add rule {rule.table.value} {rule.chain.value} {body}" cmd = f"add rule {rule.table.value} {rule.chain.value} {body}"
if rule.comment: if rule.comment:
cmd += f" comment \"{rule.comment}\"" cmd += f' comment "{rule.comment}"'
return cmd return cmd
# ---------- Endpoints -------------------------------------------------
# ------------------ Endpoints -------------------------------------------
@router.get("/options") @router.get("/options")
def get_options() -> Dict[str, Any]: def get_options() -> Dict[str, Any]:
""" """Return allowed enum choices for 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],
@@ -557,79 +548,47 @@ def get_options() -> Dict[str, Any]:
"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) -> Dict[str, Any]: def list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
""" """List rules for a table/chain (creates table/chain if missing)."""
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) 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(...)) -> Dict[str, str]: def preview_rule(rule: RuleModel = Body(...)) -> Dict[str, str]:
""" """Return nft command that would be executed for the provided rule (no-op)."""
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)) logger.error("Preview failed: %s", 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(...)) -> Dict[str, Any]: def add_rule(rule: RuleModel = Body(...)) -> Dict[str, Any]:
""" """Insert or append the rule; creates table/chain if missing."""
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) 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) -> Dict[str, Any]: def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
""" """Delete a rule by handle."""
Delete a rule by nft handle. Ensures table/chain exist before attempting delete.
"""
ensure_table_and_chain_exist(DEFAULT_FAMILY, table, chain) 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( def update_rule(handle: int, rule: RuleModel = Body(...), table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]:
handle: int, """Replace a rule by handle: delete by handle then insert at same position (if known)."""
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) 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: Optional[int] = 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 existing rule by handle
delete_rule(handle, table=table, chain=chain) delete_rule(handle, table=table, chain=chain)
# 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)