diff --git a/backend/src/api/nft_api.py b/backend/src/api/nft_api.py index 526558e..56c3369 100644 --- a/backend/src/api/nft_api.py +++ b/backend/src/api/nft_api.py @@ -1,58 +1,62 @@ +# fastapi_nft_router.py +# -*- coding: utf-8 -*- """ -FastAPI router specialized for managing nftables rules for the -'bridge' family, 'filter' table, 'forward' chain. +FastAPI router for managing nftables rules for the 'bridge' family (bridge), +default table 'mitm_tbl' and default chain 'forward'. -This enhanced version uses Python enums and narrowly-typed Pydantic -models wherever sensible so the frontend can request 'options' and -present user-friendly dropdowns. It keeps the dynamic `expr` model but -provides many structured expression types (meta, ether, ip, tcp/udp, -ct, verdict, raw) implemented as discriminated unions. +Features: +- Typed Pydantic models and enums to help the frontend populate dropdowns. +- /options endpoint to return enum choices for UI selects. +- List / preview / add / delete / update endpoints for rules. +- 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: -- Rich enums for fields (MetaKey, EtherField, IPDir, Ops, Verdict, Proto, - ConntrackState, RejectType, IcmpType, LogGroup). -- FastAPI endpoint `/options` that returns allowed enum choices for the - frontend to populate dropdowns. -- Strict Pydantic models with `kind` discriminators for expression types. -- Validation to ensure family==bridge and chain/table default to - filter/forward but still configurable if needed. -- Preview and apply endpoints unchanged in behavior but now accept - enumerated inputs which make the generated nft syntax safer and - simpler to render in the UI. -- Improved `nft_list_rules` parsing: returns rule `handle`, `position` - (1-based within chain), `comment`, and best-effort `verdict` so the - frontend can present accurate update/delete controls. +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) -Security: still run as root or with nft capabilities. """ -import logging from typing import Any, Dict, List, Optional, Union, Literal import subprocess import shutil import json +import logging from enum import Enum + from fastapi import APIRouter, HTTPException, Body from pydantic import BaseModel, Field, validator -# Router + logger -router = APIRouter() +# Router & logging ------------------------------------------------------- +router = APIRouter(prefix="/api/nft/bridge/forward", tags=["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") -# Defaults +# Defaults & local nft binary detection --------------------------------- DEFAULT_TABLE = "mitm_tbl" DEFAULT_CHAIN = "forward" DEFAULT_FAMILY = "bridge" NFT_BIN = shutil.which("nft") -# ------------------ Enums (for front-end dropdowns) ------------------ +# ------------------ Enums (for front-end dropdowns & typing) ------------ class Family(str, Enum): - bridge = "bridge" + bridge = DEFAULT_FAMILY class Table(str, Enum): - filter = "filter" - raw = "raw" + table = DEFAULT_TABLE class Chain(str, Enum): forward = "forward" @@ -101,7 +105,6 @@ class ConntrackState(str, Enum): class RejectType(str, Enum): icmp = "icmp" tcp_reset = "tcp reset" - # user may also supply raw later class IcmpType(str, Enum): dest_unreachable = "destination-unreachable" @@ -122,92 +125,237 @@ class LogGroup(int, Enum): g6 = 6 g7 = 7 -# ------------------ Pydantic expression models (discriminated unions) ------------------ +# ------------------ Pydantic expression models (discriminated unions) --- class BaseExpr(BaseModel): + """Base expression with a discriminator 'kind'.""" kind: str class Config: extra = "forbid" + class MetaExpr(BaseExpr): kind: Literal["meta"] = Field(default="meta") key: MetaKey op: Op = Op.eq value: str + class EtherExpr(BaseExpr): kind: Literal["ether"] = Field(default="ether") field: EtherField op: Op = Op.eq value: str + class IPExpr(BaseExpr): kind: Literal["ip"] = Field(default="ip") side: IPDir op: Op = Op.eq value: str + class ProtoPortExpr(BaseExpr): kind: Literal["l4"] = Field(default="l4") proto: Proto sport: Optional[str] = None dport: Optional[str] = None + class CTEexpr(BaseExpr): kind: Literal["ct"] = Field(default="ct") state: ConntrackState + class VerdictExpr(BaseExpr): kind: Literal["verdict"] = Field(default="verdict") verdict: Verdict + class RejectExpr(BaseExpr): kind: Literal["reject"] = Field(default="reject") reject_type: RejectType icmp_type: Optional[IcmpType] = None + class LogExpr(BaseExpr): kind: Literal["log"] = Field(default="log") prefix: Optional[str] = None group: Optional[LogGroup] = None + class RawExpr(BaseExpr): kind: Literal["raw"] = Field(default="raw") 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): + """Model for creating/inserting rules through the API.""" family: Family = Family.bridge - table: Table = Table.filter + table: Table = Table.table chain: Chain = Chain.forward expr: List[Expr] = Field(default_factory=list) 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 @validator("family") - def only_bridge(cls, v): + def only_bridge(cls, v: Family) -> Family: if v != Family.bridge: raise ValueError("This router only manages family 'bridge'") 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: + logger.error("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() - 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: - 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 "") + 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 + - add chain
{ type filter hook priority 0; policy accept; } + for chain in (input, forward, output). For other chain names a simple chain is created + without hook: `add 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) 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()}") results: List[Dict[str, Any]] = [] @@ -218,7 +366,6 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A items = [] for item in items: - # skip table/chain metadata entries (we only need rule entries) if "rule" not in item: continue @@ -232,37 +379,31 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A position = counters[key] handle = r.get("handle") - exprs = r.get("expr", []) # original expression list from nft JSON - comment = None - verdict = None - verdict_details = None + exprs = r.get("expr", []) # original expression list + comment: Optional[str] = None + verdict: Optional[str] = None + verdict_details: Optional[Any] = None # scan expressions to extract comment and verdict/action for expr in exprs: if not isinstance(expr, dict): continue - # comment can appear as {"comment":"text"} or {"comment": {"text": "..."}} depending on nft json variant if "comment" in expr: - # handle both simple and nested forms c = expr.get("comment") if isinstance(c, str): comment = c elif isinstance(c, dict): - # some representations: {"comment": {"text": "..."}} or {"comment": {"str": "..."}} 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: v = expr["verdict"] - # v often looks like {"accept": None} or {"drop": None} or {"reject": {...}} if isinstance(v, dict): - # take the first key as the action k = next(iter(v.keys()), None) verdict = k verdict_details = v.get(k) else: verdict = str(v) - # older/alternate forms if "drop" in expr and verdict is None: verdict = "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_details = expr.get("reject") - results.append({ - "family": family, - "table": table_name, - "chain": chain_name, - "handle": handle, - "position": position, - "comment": comment, - "verdict": verdict, # e.g. "accept", "drop", "reject", or None - "verdict_details": verdict_details, # raw details for rejects/other actions - "exprs": exprs, # original expression list (renamed, clearer) - "nft_rule": r, # the original nft JSON dict for this rule - }) + results.append( + { + "family": family, + "table": table_name, + "chain": chain_name, + "handle": handle, + "position": position, + "comment": comment, + "verdict": verdict, + "verdict_details": verdict_details, + "exprs": exprs, + "nft_rule": r, + } + ) # filter by requested table/chain if provided if table or chain: @@ -292,9 +435,8 @@ def nft_list_rules(table: str = "filter", chain: str = "forward") -> Dict[str, A return {"rules": results} - 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): val = e.value key = e.key.value @@ -318,8 +460,7 @@ def expr_to_nft_snippet(e: Expr) -> str: if e.reject_type == RejectType.icmp: t = f"reject with icmp type {e.icmp_type.value}" if e.icmp_type else "reject" return t - else: - return e.reject_type.value + return e.reject_type.value if isinstance(e, LogExpr): parts = ["log"] if e.prefix: @@ -333,6 +474,7 @@ def expr_to_nft_snippet(e: Expr) -> 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] body = " ".join(s for s in expr_snippets if s) if rule.position is not None: @@ -344,21 +486,13 @@ def rule_to_nft_cmd(rule: RuleModel) -> str: return cmd -def run_nft_cmd(cmd: str) -> Dict[str, Any]: - 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 ------------------ - +# ------------------ Endpoints ------------------------------------------- @router.get("/options") -def get_options(): - """Return available enum choices for the frontend dropdowns.""" +def get_options() -> Dict[str, Any]: + """ + Return allowed enum choices for the frontend dropdowns. + The frontend should call this once and cache results to populate selects. + """ return { "family": [f.value for f in Family], "table": [t.value for t in Table], @@ -375,42 +509,79 @@ def get_options(): "log_groups": [int(g.value) for g in LogGroup], } + @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) + @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: cmd = rule_to_nft_cmd(rule) except Exception as e: + logger.error("Preview building failed: %s", str(e)) raise HTTPException(status_code=400, detail=str(e)) return {"cmd": cmd} + @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) return run_nft_cmd(cmd) + @router.delete("/rules/{handle}") -def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN): - ensure_nft_available() +def delete_rule(handle: int, table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN) -> Dict[str, Any]: + """ + 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}" return run_nft_cmd(cmd) + @router.put("/rules/{handle}") -def update_rule(handle: int, rule: RuleModel = Body(...), table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN): - # Delete by handle then insert at the same position (if available via list parse) - # We attempt to find the position of the handle so the replacement keeps position. +def update_rule( + handle: int, + 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) - position = None - for r in rules_info.get('rules', []): - if r.get('handle') == handle: - position = r.get('position') + position: Optional[int] = None + for r in rules_info.get("rules", []): + if r.get("handle") == handle: + position = r.get("position") break - # delete by handle + + # delete existing rule by handle 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: rule.position = position return add_rule(rule)