This commit is contained in:
410
backend/src/api/nft_api.py
Normal file
410
backend/src/api/nft_api.py
Normal file
@@ -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)
|
||||||
@@ -5,12 +5,14 @@ import os
|
|||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
|
||||||
from src.api import packet_api
|
|
||||||
from src.utilities.packet_broadcaster import PacketBroadcaster
|
from src.utilities.packet_broadcaster import PacketBroadcaster
|
||||||
import src.shared_objects as shared_objects
|
import src.shared_objects as shared_objects
|
||||||
from src.utilities.database import DatabasePool
|
from src.utilities.database import DatabasePool
|
||||||
import src.api.network_api as network_api
|
import src.api.network_api as network_api
|
||||||
import src.api.sniffer_api as sniffer_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
|
import src.api.nftables_api as nftables_api
|
||||||
|
|
||||||
# ---- Config -----------------------------------------------------------
|
# ---- Config -----------------------------------------------------------
|
||||||
@@ -134,4 +136,5 @@ def versions():
|
|||||||
app.include_router(network_api.router, prefix="/network", tags=["network"])
|
app.include_router(network_api.router, prefix="/network", tags=["network"])
|
||||||
app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"])
|
app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"])
|
||||||
app.include_router(packet_api.router, prefix="/packets", tags=["packets"])
|
app.include_router(packet_api.router, prefix="/packets", tags=["packets"])
|
||||||
app.include_router(nftables_api.router, prefix="/nftables", tags=["nftables"])
|
app.include_router(nftables_api.router, prefix="/nftables", tags=["nftables"])
|
||||||
|
app.include_router(nft_api.router, prefix="/nft", tags=["nft"])
|
||||||
Reference in New Issue
Block a user