diff --git a/backend/src/api/nft_api.py b/backend/src/api/nft_api.py new file mode 100644 index 0000000..0118e6a --- /dev/null +++ b/backend/src/api/nft_api.py @@ -0,0 +1,410 @@ +""" +FastAPI router specialized for managing nftables rules for the +'bridge' family, 'filter' table, 'forward' chain. + +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 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: 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 +from enum import Enum +from fastapi import APIRouter, HTTPException, Body +from pydantic import BaseModel, Field, validator + +# Router + logger +router = APIRouter() +logger = logging.getLogger("nftables") +logger.debug("nftables router module loaded") + +# Defaults +DEFAULT_TABLE = "mitm_tbl" +DEFAULT_CHAIN = "forward" +DEFAULT_FAMILY = "bridge" + +NFT_BIN = shutil.which("nft") + +# ------------------ Enums (for front-end dropdowns) ------------------ +class Family(str, Enum): + bridge = "bridge" + +class Table(str, Enum): + filter = "filter" + raw = "raw" + +class Chain(str, Enum): + forward = "forward" + input = "input" + output = "output" + +class MetaKey(str, Enum): + iifname = "iifname" + oifname = "oifname" + iif = "iif" + oif = "oif" + prio = "prio" + +class EtherField(str, Enum): + saddr = "saddr" + daddr = "daddr" + +class IPDir(str, Enum): + saddr = "saddr" + daddr = "daddr" + +class Op(str, Enum): + eq = "==" + neq = "!=" + lt = "<" + gt = ">" + contains = "in" + +class Verdict(str, Enum): + accept = "accept" + drop = "drop" + reject = "reject" + continue_ = "continue" + +class Proto(str, Enum): + tcp = "tcp" + udp = "udp" + icmp = "icmp" + +class ConntrackState(str, Enum): + new = "new" + established = "established" + related = "related" + invalid = "invalid" + +class RejectType(str, Enum): + icmp = "icmp" + tcp_reset = "tcp reset" + # user may also supply raw later + +class IcmpType(str, Enum): + dest_unreachable = "destination-unreachable" + time_exceeded = "time-exceeded" + echo_reply = "echo-reply" + echo_request = "echo-request" + port_unreachable = "port-unreachable" + host_unreachable = "host-unreachable" + fragmentation_needed = "fragmentation-needed" + +class LogGroup(int, Enum): + g0 = 0 + g1 = 1 + g2 = 2 + g3 = 3 + g4 = 4 + g5 = 5 + g6 = 6 + g7 = 7 + +# ------------------ Pydantic expression models (discriminated unions) ------------------ +class BaseExpr(BaseModel): + kind: str + + class Config: + extra = "forbid" + +class MetaExpr(BaseExpr): + kind: Literal["meta"] = Field("meta", const=True) + key: MetaKey + op: Op = Op.eq + value: str + +class EtherExpr(BaseExpr): + kind: Literal["ether"] = Field("ether", const=True) + field: EtherField + op: Op = Op.eq + value: str + +class IPExpr(BaseExpr): + kind: Literal["ip"] = Field("ip", const=True) + side: IPDir + op: Op = Op.eq + value: str + +class ProtoPortExpr(BaseExpr): + kind: Literal["l4"] = Field("l4", const=True) + proto: Proto + sport: Optional[str] = None + dport: Optional[str] = None + +class CTEexpr(BaseExpr): + kind: Literal["ct"] = Field("ct", const=True) + state: ConntrackState + +class VerdictExpr(BaseExpr): + kind: Literal["verdict"] = Field("verdict", const=True) + verdict: Verdict + +class RejectExpr(BaseExpr): + kind: Literal["reject"] = Field("reject", const=True) + reject_type: RejectType + icmp_type: Optional[IcmpType] = None + +class LogExpr(BaseExpr): + kind: Literal["log"] = Field("log", const=True) + prefix: Optional[str] = None + group: Optional[LogGroup] = None + +class RawExpr(BaseExpr): + kind: Literal["raw"] = Field("raw", const=True) + snippet: str + +# Combine into a discriminated union +Expr = Union[MetaExpr, EtherExpr, IPExpr, ProtoPortExpr, CTEexpr, VerdictExpr, RejectExpr, LogExpr, RawExpr] + +# ------------------ Rule model ------------------ +class RuleModel(BaseModel): + family: Family = Family.bridge + table: Table = Table.filter + chain: Chain = Chain.forward + expr: List[Expr] = Field(default_factory=list) + comment: Optional[str] = None + position: Optional[int] = None # position to insert (1-based) + handle: Optional[int] = None + + @validator("family") + def only_bridge(cls, v): + if v != Family.bridge: + raise ValueError("This router only manages family 'bridge'") + return v + +# ------------------ Helpers ------------------ + +def ensure_nft_available(): + if not NFT_BIN: + 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]: + """Return nft rules parsed from `nft --json list ruleset`. + + For each rule we return a dict with at least: + - family, table, chain + - handle (if present) + - position (1-based index within chain) + - comment (if present in expr) + - verdict (best-effort string like 'accept'/'drop'/'reject') + - raw_exprs: the original expr list from nft JSON + + This makes it easier for the frontend to show position and edit + or replace a specific rule by handle/position. + """ + ensure_nft_available() + cmd = [NFT_BIN, "--json", "list", "ruleset"] + try: + out = subprocess.check_output(cmd, stderr=subprocess.PIPE) + parsed = json.loads(out) + except subprocess.CalledProcessError as e: + raise HTTPException(status_code=500, detail=f"nft failed: {e.stderr.decode()}") + + results: List[Dict[str, Any]] = [] + # counters per chain to compute position + counters: Dict[str, int] = {} + + # The JSON usually contains a top-level dict with key 'nftables' -> list + items = parsed.get("nftables") if isinstance(parsed, dict) else parsed + if not isinstance(items, list): + items = [] + + for item in items: + # Keep a simple mapping of current table/chain context + if 'table' in item: + tbl = item['table'] + # nothing to do for context; continue + continue + if 'chain' in item: + ch = item['chain'] + # continue; chain metadata present + continue + if 'rule' in item: + r = item['rule'] + family = r.get('family') + table_name = r.get('table') + chain_name = r.get('chain') + key = f"{family}:{table_name}:{chain_name}" + counters.setdefault(key, 0) + counters[key] += 1 + position = counters[key] + handle = r.get('handle') + # extract comment and verdict best-effort + raw_exprs = r.get('expr', []) + comment = None + verdict = None + for expr in raw_exprs: + if isinstance(expr, dict): + if 'comment' in expr: + comment = expr.get('comment') + if 'verdict' in expr: + # verdict might be dict like {'verdict': {'kind': 'accept'}} + v = expr['verdict'] + if isinstance(v, dict): + verdict = list(v.keys())[0] + else: + verdict = str(v) + # some JSON formats have 'match' or 'payload' etc. look for 'type' fields + if 'reject' in expr: + verdict = 'reject' + results.append({ + 'family': family, + 'table': table_name, + 'chain': chain_name, + 'handle': handle, + 'position': position, + 'comment': comment, + 'verdict': verdict, + 'raw_exprs': raw_exprs, + 'raw_rule': r, + }) + # Optionally filter by table/chain if requested + if table or chain: + results = [x for x in results if x['table'] == table and x['chain'] == chain] + return {"rules": results} + + +def expr_to_nft_snippet(e: Expr) -> str: + # runtime dispatch via model type + if isinstance(e, MetaExpr): + val = e.value + key = e.key.value + return f"meta {key} {e.op.value} {val}" + if isinstance(e, EtherExpr): + return f"ether {e.field.value} {e.op.value} {e.value}" + if isinstance(e, IPExpr): + return f"ip {e.side.value} {e.op.value} {e.value}" + if isinstance(e, ProtoPortExpr): + parts = [e.proto.value] + if e.sport: + parts.append(f"sport {e.sport}") + if e.dport: + parts.append(f"dport {e.dport}") + return " ".join(parts) + if isinstance(e, CTEexpr): + return f"ct state {e.state.value}" + if isinstance(e, VerdictExpr): + return e.verdict.value if e.verdict != Verdict.continue_ else "continue" + if isinstance(e, RejectExpr): + 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 + if isinstance(e, LogExpr): + parts = ["log"] + if e.prefix: + parts.append(f"prefix \"{e.prefix}\"") + if e.group is not None: + parts.append(f"group {int(e.group)}") + return " ".join(parts) + if isinstance(e, RawExpr): + return e.snippet + raise ValueError("Unsupported expression type") + + +def rule_to_nft_cmd(rule: RuleModel) -> str: + 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: + cmd = f"insert rule {rule.table.value} {rule.chain.value} position {rule.position} {body}" + else: + cmd = f"add rule {rule.table.value} {rule.chain.value} {body}" + if rule.comment: + cmd += f" comment \"{rule.comment}\"" + 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 ------------------ + +@router.get("/options") +def get_options(): + """Return available enum choices for the frontend dropdowns.""" + return { + "family": [f.value for f in Family], + "table": [t.value for t in Table], + "chain": [c.value for c in Chain], + "meta_keys": [m.value for m in MetaKey], + "ether_fields": [e.value for e in EtherField], + "ip_dirs": [d.value for d in IPDir], + "ops": [o.value for o in Op], + "verdicts": [v.value for v in Verdict], + "protocols": [p.value for p in Proto], + "ct_states": [s.value for s in ConntrackState], + "reject_types": [r.value for r in RejectType], + "icmp_types": [i.value for i in IcmpType], + "log_groups": [int(g.value) for g in LogGroup], + } + +@router.get("/rules") +def list_rules(table: str = DEFAULT_TABLE, chain: str = DEFAULT_CHAIN): + return nft_list_rules(table=table, chain=chain) + +@router.post("/rules/preview") +def preview_rule(rule: RuleModel = Body(...)): + try: + cmd = rule_to_nft_cmd(rule) + except Exception as e: + raise HTTPException(status_code=400, detail=str(e)) + return {"cmd": cmd} + +@router.post("/rules") +def add_rule(rule: RuleModel = Body(...)): + 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() + 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. + 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') + break + # delete by handle + delete_rule(handle, table=table, chain=chain) + # if we found position, insert at that position; otherwise append + if position is not None: + rule.position = position + return add_rule(rule) diff --git a/backend/src/main.py b/backend/src/main.py index 00046c8..e552d94 100644 --- a/backend/src/main.py +++ b/backend/src/main.py @@ -5,12 +5,14 @@ import os from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware -from src.api import packet_api + from src.utilities.packet_broadcaster import PacketBroadcaster import src.shared_objects as shared_objects from src.utilities.database import DatabasePool import src.api.network_api as network_api import src.api.sniffer_api as sniffer_api +from src.api import nft_api +from src.api import packet_api import src.api.nftables_api as nftables_api # ---- Config ----------------------------------------------------------- @@ -134,4 +136,5 @@ def versions(): app.include_router(network_api.router, prefix="/network", tags=["network"]) app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"]) app.include_router(packet_api.router, prefix="/packets", tags=["packets"]) -app.include_router(nftables_api.router, prefix="/nftables", tags=["nftables"]) \ No newline at end of file +app.include_router(nftables_api.router, prefix="/nftables", tags=["nftables"]) +app.include_router(nft_api.router, prefix="/nft", tags=["nft"]) \ No newline at end of file