file structure and comments unified
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s

This commit is contained in:
2026-03-06 19:03:26 +01:00
parent 7b9a7d3a4b
commit b60a1d3118
20 changed files with 697 additions and 1823 deletions

View File

@@ -1,13 +1,20 @@
"""Netplan schema models used by bridge/network configuration APIs."""
from typing import Dict, List, Optional
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import List, Dict, Optional
class Nameservers(BaseModel): class Nameservers(BaseModel):
"""DNS nameserver configuration."""
addresses: List[str] = Field(default_factory=list) addresses: List[str] = Field(default_factory=list)
search: List[str] = Field(default_factory=list) search: List[str] = Field(default_factory=list)
class EthernetConfig(BaseModel): class EthernetConfig(BaseModel):
"""Netplan ethernet interface configuration."""
dhcp4: Optional[bool] = None dhcp4: Optional[bool] = None
dhcp6: Optional[bool] = None dhcp6: Optional[bool] = None
addresses: Optional[List[str]] = None addresses: Optional[List[str]] = None
@@ -18,44 +25,23 @@ class EthernetConfig(BaseModel):
class BridgeConfig(BaseModel): class BridgeConfig(BaseModel):
interfaces: List[str] = Field(default_factory=list) # ["eth1", "eth2"] """Netplan bridge configuration."""
interfaces: List[str] = Field(default_factory=list)
dhcp4: Optional[bool] = None dhcp4: Optional[bool] = None
dhcp6: Optional[bool] = None dhcp6: Optional[bool] = None
addresses: Optional[List[str]] = None addresses: Optional[List[str]] = None
gateway4: Optional[str] = None gateway4: Optional[str] = None
gateway6: Optional[str] = None gateway6: Optional[str] = None
nameservers: Optional[Nameservers] = None nameservers: Optional[Nameservers] = None
parameters: Optional[dict] = None # allows spanning-tree, port-priority, forward-delay, etc. parameters: Optional[dict] = None
optional: Optional[bool] = None optional: Optional[bool] = None
class NetworkConfig(BaseModel): class NetworkConfig(BaseModel):
"""Top-level Netplan network object."""
version: int = 2 version: int = 2
renderer: Optional[str] = "networkd" renderer: Optional[str] = "networkd"
ethernets: Dict[str, EthernetConfig] = Field(default_factory=dict) ethernets: Dict[str, EthernetConfig] = Field(default_factory=dict)
bridges: Dict[str, BridgeConfig] = Field(default_factory=dict) bridges: Dict[str, BridgeConfig] = Field(default_factory=dict)
'''Example usage:
{
"version": 2,
"renderer": "networkd",
"ethernets": {
"eth0": {
"dhcp4": false,
"addresses": ["192.168.10.20/24"],
"gateway4": "192.168.10.1",
"nameservers": {
"addresses": ["1.1.1.1", "8.8.8.8"]
}
},
"eth1": {},
"eth2": {}
},
"bridges": {
"br0": {
"interfaces": ["eth1", "eth2"],
"dhcp4": true
}
}
}'''

View File

@@ -1,12 +1,16 @@
# src/models/packet.py """Pydantic model for packet rows returned by the backend."""
from datetime import datetime
from pydantic import BaseModel, Field, IPvAnyAddress from datetime import datetime
from typing import Optional, Union from typing import Optional, Union
from pydantic import BaseModel, Field, IPvAnyAddress
class PacketDBModel(BaseModel): class PacketDBModel(BaseModel):
"""Normalized packet representation used across DB and API layers."""
id: Union[int, str] id: Union[int, str]
timestamp: datetime = Field(..., description="ISO timestamp") timestamp: datetime = Field(..., description="Packet timestamp in ISO format.")
iface: str iface: str
src_mac: Optional[str] = None src_mac: Optional[str] = None
dst_mac: Optional[str] = None dst_mac: Optional[str] = None
@@ -18,7 +22,7 @@ class PacketDBModel(BaseModel):
dst_port: Optional[int] = None dst_port: Optional[int] = None
vlan_id: Optional[int] = None vlan_id: Optional[int] = None
length: Optional[int] = None length: Optional[int] = None
raw_b64: Optional[str] = Field(None, description="Base64-encoded raw bytes") raw_b64: Optional[str] = Field(None, description="Base64-encoded packet bytes.")
direction: Optional[str] = None direction: Optional[str] = None
packets: Optional[int] = None packets: Optional[int] = None
@@ -40,4 +44,4 @@ class PacketDBModel(BaseModel):
"length": 128, "length": 128,
"raw_b64": "BASE64...", "raw_b64": "BASE64...",
} }
} }

View File

@@ -1,122 +1,95 @@
"""Network inspection and bridge management endpoints."""
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import List, Optional
from pyroute2 import IPRoute, NDB from pyroute2 import IPRoute, NDB
router = APIRouter() router = APIRouter()
# Globals for lazy initialization
ip: IPRoute | None = None ip: IPRoute | None = None
ndb: NDB | None = None ndb: NDB | None = None
# ------------------------------
# Pydantic models
# ------------------------------
class InterfaceAddress(BaseModel): class InterfaceAddress(BaseModel):
""" """IP address assigned to an interface."""
Represents an IP address assigned to a network interface.
"""
family: str = Field(..., description="IP family: 'ipv4' or 'ipv6'.") family: str = Field(..., description="IP family: 'ipv4' or 'ipv6'.")
address: str = Field(..., description="The IP address assigned to the interface.") address: str = Field(..., description="IP address.")
prefixlen: int = Field(..., description="Subnet prefix length (e.g., 24 for 255.255.255.0).") prefixlen: int = Field(..., description="Subnet prefix length.")
class InterfaceInfo(BaseModel): class InterfaceInfo(BaseModel):
""" """Interface with link metadata and assigned addresses."""
Represents a network interface with all its properties.
""" ifindex: int = Field(..., description="Kernel interface index.")
ifindex: int = Field(..., description="Interface index (unique identifier assigned by the kernel).") name: str = Field(..., description="Interface name.")
name: str = Field(..., description="Interface name (e.g., 'eth0', 'enp38s0').") state: str = Field(..., description="Operational state.")
state: str = Field(..., description="Operational state (e.g., 'UP', 'DOWN', 'UNKNOWN').") mac: Optional[str] = Field(None, description="MAC address.")
mac: Optional[str] = Field(None, description="MAC address of the interface, if applicable.") mtu: int = Field(..., description="Maximum transmission unit.")
mtu: int = Field(..., description="Maximum Transmission Unit for the interface.") flags: List[str] = Field(..., description="Decoded interface flags.")
flags: List[str] = Field(..., description="List of interface flags (e.g., ['BROADCAST', 'MULTICAST']).") addresses: List[InterfaceAddress] = Field(..., description="Assigned IP addresses.")
addresses: List[InterfaceAddress] = Field(..., description="List of IP addresses assigned to the interface.")
class RouteInfo(BaseModel): class RouteInfo(BaseModel):
""" """Single routing table entry."""
Represents a single routing table entry.
"""
dst: Optional[str] = Field( dst: Optional[str] = Field(None, description="Destination CIDR; null means default route.")
None, description="Destination network in CIDR notation (e.g., '192.168.1.0/24'). None means default route." gateway: Optional[str] = Field(None, description="Next-hop gateway.")
) prefsrc: Optional[str] = Field(None, description="Preferred source IP.")
oif: Optional[int] = Field(None, description="Output interface index.")
ifname: Optional[str] = Field(None, description="Output interface name.")
table: int = Field(..., description="Route table ID.")
proto: Optional[int] = Field(None, description="Route protocol code.")
scope: Optional[int] = Field(None, description="Route scope code.")
type: Optional[int] = Field(None, description="Route type code.")
gateway: Optional[str] = Field(
None, description="Next-hop gateway IP address for this route. None if the route is directly connected."
)
prefsrc: Optional[str] = Field(
None, description="Preferred source IP to use when sending packets via this route."
)
oif: Optional[int] = Field(
None, description="Output interface index (ifindex) for this route. Can be used to look up the interface name."
)
ifname: Optional[str] = Field(
None, description="Name of the interface corresponding to `oif` (e.g., 'eth0')."
)
table: int = Field(
..., description="Routing table ID (e.g., 254 = main, 255 = local)."
)
proto: Optional[int] = Field(
None,
description="Protocol of the route (numeric Linux codes, e.g., 2=kernel, 16=static)."
)
scope: Optional[int] = Field(
None,
description="Scope of the route: 0=global, 253=link, 254=host, 255=nowhere."
)
type: Optional[int] = Field(
None,
description="Type of the route (numeric code): 1=unicast, 2=local, 3=broadcast, 5=multicast."
)
class BridgeInterfaceInfo(BaseModel): class BridgeInterfaceInfo(BaseModel):
""" """Interface that belongs to a bridge."""
Represents a network interface which is a member of an bridge.
""" ifindex: int = Field(..., description="Interface index.")
ifindex: int = Field(..., description="Interface index of a bridge member") ifname: str = Field(..., description="Interface name.")
ifname: str = Field(..., description="Interface name of a bridge member") state: Optional[str] = Field(None, description="Operational state.")
state: Optional[str] = Field(None, description="Operational state of the interface") mtu: Optional[int] = Field(None, description="Interface MTU.")
mtu: Optional[int] = Field(None, description="MTU of the interface")
class BridgeInfo(BaseModel): class BridgeInfo(BaseModel):
""" """Bridge interface with member information."""
Represents a network bridge interface with all its properties.
""" ifindex: int = Field(..., description="Bridge index.")
ifindex: int = Field(..., description="Interface index of the bridge") ifname: str = Field(..., description="Bridge name.")
ifname: str = Field(..., description="Bridge interface name") state: Optional[str] = Field(None, description="Bridge state.")
state: Optional[str] = Field(None, description="Operational state of the bridge") mtu: Optional[int] = Field(None, description="Bridge MTU.")
mtu: Optional[int] = Field(None, description="MTU of the bridge") stp_state: Optional[int] = Field(None, description="Spanning tree state.")
stp_state: Optional[int] = Field(None, description="STP (Spanning Tree Protocol) state of the bridge") members: List[BridgeInterfaceInfo] = Field(default_factory=list, description="Bridge members.")
members: List[BridgeInterfaceInfo] = Field(default_factory=list, description="List of member interfaces of the bridge")
class BridgeCreateRequest(BaseModel): class BridgeCreateRequest(BaseModel):
"""Payload for creating a bridge and attaching interfaces."""
name: str name: str
interfaces: List[str] interfaces: List[str]
class BridgeRemoveRequest(BaseModel): class BridgeRemoveRequest(BaseModel):
name: str """Payload for removing a bridge."""
# ------------------------------
# Lazy Init Functions
# ------------------------------
def init_network_api(): name: str
def init_network_api() -> None:
"""Initialize lazy pyroute2 clients."""
global ip, ndb global ip, ndb
if ip is None: if ip is None:
ip = IPRoute() ip = IPRoute()
if ndb is None: if ndb is None:
ndb = NDB() ndb = NDB()
def shutdown_network_api():
def shutdown_network_api() -> None:
"""Close pyroute2 clients if they were initialized."""
global ip, ndb global ip, ndb
if ip: if ip:
ip.close() ip.close()
@@ -125,37 +98,38 @@ def shutdown_network_api():
ndb.close() ndb.close()
ndb = None ndb = None
def get_iproute():
def get_iproute() -> IPRoute:
"""Dependency provider for the shared IPRoute instance."""
if ip is None: if ip is None:
init_network_api() init_network_api()
return ip return ip
def get_ndb():
def get_ndb() -> NDB:
"""Dependency provider for the shared NDB instance."""
if ndb is None: if ndb is None:
init_network_api() init_network_api()
return ndb return ndb
# ------------------------------
# Utility functions
# ------------------------------
def parse_addresses(addrs): def parse_addresses(addrs: list[dict]) -> list[InterfaceAddress]:
res = [] """Convert pyroute2 address rows into `InterfaceAddress` models."""
for a in addrs: result: list[InterfaceAddress] = []
family = "ipv4" if a.get("family") == 2 else "ipv6" for addr in addrs:
res.append( family = "ipv4" if addr.get("family") == 2 else "ipv6"
result.append(
InterfaceAddress( InterfaceAddress(
family=family, family=family,
address=a.get("address"), address=addr.get("address"),
prefixlen=a.get("prefixlen"), prefixlen=addr.get("prefixlen"),
) )
) )
return res return result
def parse_flags(flags_int: int) -> list[str]: def parse_flags(flags_int: int) -> list[str]:
""" """Decode Linux interface flag bitset to names."""
Converts the integer flags from pyroute2 to human-readable list of strings.
"""
flags_map = { flags_map = {
0x1: "UP", 0x1: "UP",
0x2: "BROADCAST", 0x2: "BROADCAST",
@@ -177,37 +151,33 @@ def parse_flags(flags_int: int) -> list[str]:
0x20000: "DORMANT", 0x20000: "DORMANT",
0x40000: "ECHO", 0x40000: "ECHO",
} }
result = [] return [name for bit, name in flags_map.items() if flags_int & bit]
for bit, name in flags_map.items():
if flags_int & bit:
result.append(name)
return result
def iface_index(name: str, ip: IPRoute) -> int:
idx = ip.link_lookup(ifname=name) def iface_index(name: str, ip_route: IPRoute) -> int:
"""Return interface index for a given interface name."""
idx = ip_route.link_lookup(ifname=name)
if not idx: if not idx:
raise HTTPException(status_code=404, detail=f"Interface {name} not found") raise HTTPException(status_code=404, detail=f"Interface {name} not found")
return idx[0] return idx[0]
def bridge_exists(name: str, ip: IPRoute) -> bool: def bridge_exists(name: str, ip_route: IPRoute) -> bool:
return bool(ip.link_lookup(ifname=name)) """Check whether a bridge/device with the given name exists."""
return bool(ip_route.link_lookup(ifname=name))
# ------------------------------
# Endpoints
# ------------------------------
@router.get("/interfaces", response_model=List[InterfaceInfo]) @router.get("/interfaces", response_model=List[InterfaceInfo])
def get_interfaces(ip: IPRoute = Depends(get_iproute)): def get_interfaces(ip: IPRoute = Depends(get_iproute)) -> List[InterfaceInfo]:
result = [] """List host interfaces with addresses and decoded flags."""
result: list[InterfaceInfo] = []
links = ip.get_links() links = ip.get_links()
addresses = ip.get_addr() addresses = ip.get_addr()
addr_map = {} addr_map: dict[int, list] = {}
for a in addresses: for addr in addresses:
ifindex = a.get("index") ifindex = addr.get("index")
addr_map.setdefault(ifindex, []).append(a) addr_map.setdefault(ifindex, []).append(addr)
for link in links: for link in links:
attrs = dict(link["attrs"]) attrs = dict(link["attrs"])
@@ -227,55 +197,53 @@ def get_interfaces(ip: IPRoute = Depends(get_iproute)):
) )
return result return result
@router.get("/routes", response_model=List[RouteInfo]) @router.get("/routes", response_model=List[RouteInfo])
def get_routes(ip: IPRoute = Depends(get_iproute)): def get_routes(ip: IPRoute = Depends(get_iproute)) -> List[RouteInfo]:
routes = [] """List routes from the kernel routing tables."""
for r in ip.get_routes(): routes: list[RouteInfo] = []
attrs = dict(r["attrs"]) for route in ip.get_routes():
attrs = dict(route["attrs"])
dst = attrs.get("RTA_DST") dst = attrs.get("RTA_DST")
gateway = attrs.get("RTA_GATEWAY") gateway = attrs.get("RTA_GATEWAY")
prefsrc = attrs.get("RTA_PREFSRC") prefsrc = attrs.get("RTA_PREFSRC")
oif = r.get("oif") oif = route.get("oif")
ifname = None ifname = None
if oif is not None: if oif is not None:
# translate ifindex → name
link = ip.get_links(oif)[0] link = ip.get_links(oif)[0]
ifname = dict(link["attrs"]).get("IFLA_IFNAME") ifname = dict(link["attrs"]).get("IFLA_IFNAME")
routes.append( routes.append(
RouteInfo( RouteInfo(
dst=f"{dst}/{r.get('dst_len')}" if dst else None, dst=f"{dst}/{route.get('dst_len')}" if dst else None,
gateway=gateway, gateway=gateway,
prefsrc=prefsrc, prefsrc=prefsrc,
oif=oif, oif=oif,
ifname=ifname, ifname=ifname,
table=r.get("table", 254), table=route.get("table", 254),
proto=r.get("proto"), proto=route.get("proto"),
scope=r.get("scope"), scope=route.get("scope"),
type=r.get("type"), type=route.get("type"),
) )
) )
return routes return routes
@router.get("/links", response_model=List[InterfaceInfo]) @router.get("/links", response_model=List[InterfaceInfo])
def get_raw_links(ip: IPRoute = Depends(get_iproute)): def get_raw_links(ip: IPRoute = Depends(get_iproute)) -> List[InterfaceInfo]:
""" """List links in a normalized structure for UI consumers."""
Returns all interfaces in a clean Pydantic format. result: list[InterfaceInfo] = []
This is similar to /interfaces but avoids additional processing if needed.
"""
result = []
links = ip.get_links() links = ip.get_links()
addresses = ip.get_addr() addresses = ip.get_addr()
# group addresses by interface index addr_map: dict[int, list] = {}
addr_map = {} for addr in addresses:
for a in addresses: ifindex = addr.get("index")
ifindex = a.get("index") addr_map.setdefault(ifindex, []).append(addr)
addr_map.setdefault(ifindex, []).append(a)
for link in links: for link in links:
attrs = dict(link.get("attrs", [])) # convert list of tuples to dict attrs = dict(link.get("attrs", []))
ifindex = link["index"] ifindex = link["index"]
addrs = addr_map.get(ifindex, []) addrs = addr_map.get(ifindex, [])
@@ -286,116 +254,95 @@ def get_raw_links(ip: IPRoute = Depends(get_iproute)):
state=attrs.get("IFLA_OPERSTATE", "unknown"), state=attrs.get("IFLA_OPERSTATE", "unknown"),
mac=attrs.get("IFLA_ADDRESS"), mac=attrs.get("IFLA_ADDRESS"),
mtu=attrs.get("IFLA_MTU", 0), mtu=attrs.get("IFLA_MTU", 0),
flags=[], # latest pyroute2 removed ifi_flags, leave empty flags=[],
addresses=parse_addresses(addrs), addresses=parse_addresses(addrs),
) )
) )
return result return result
@router.get("/bridges", response_model=List[BridgeInfo]) @router.get("/bridges", response_model=List[BridgeInfo])
def get_bridges(): def get_bridges() -> List[BridgeInfo]:
""" """List all bridges and their current member interfaces."""
Get all bridge interfaces on the system, including their member interfaces. bridges_list: list[BridgeInfo] = []
Returns detailed information:
- Bridge index, name, state, MTU
- STP state
- Member interfaces with index, name, state, and MTU
"""
bridges_list: List[BridgeInfo] = []
with NDB() as ndb: with NDB() as ndb_ctx:
for br in ndb.interfaces: for bridge in ndb_ctx.interfaces:
# Only bridges if getattr(bridge, "kind", None) != "bridge":
if getattr(br, "kind", None) == "bridge": continue
members: List[BridgeInterfaceInfo] = []
# Find member interfaces members: list[BridgeInterfaceInfo] = []
for iface in ndb.interfaces: for iface in ndb_ctx.interfaces:
if getattr(iface, "master", None) == br.index: if getattr(iface, "master", None) == bridge.index:
members.append( members.append(
BridgeInterfaceInfo( BridgeInterfaceInfo(
ifindex=iface.index, ifindex=iface.index,
ifname=iface.ifname, ifname=iface.ifname,
state=getattr(iface, "operstate", None), state=getattr(iface, "operstate", None),
mtu=getattr(iface, "mtu", None) mtu=getattr(iface, "mtu", None),
)
) )
bridges_list.append(
BridgeInfo(
ifindex=br.index,
ifname=br.ifname,
state=getattr(br, "operstate", None),
mtu=getattr(br, "mtu", None),
stp_state=getattr(br, "stp_state", None),
members=members
) )
bridges_list.append(
BridgeInfo(
ifindex=bridge.index,
ifname=bridge.ifname,
state=getattr(bridge, "operstate", None),
mtu=getattr(bridge, "mtu", None),
stp_state=getattr(bridge, "stp_state", None),
members=members,
) )
)
return bridges_list return bridges_list
@router.get("/full-state") @router.get("/full-state")
def full_state( def full_state(ip: IPRoute = Depends(get_iproute)) -> dict:
ip: IPRoute = Depends(get_iproute), """Return interfaces, routes, and bridges in one response."""
):
"""
Returns the full network state:
- Interfaces with IP addresses and flags
- Routes
- Bridges with member interfaces
"""
return { return {
"interfaces": get_interfaces(ip), "interfaces": get_interfaces(ip),
"routes": get_routes(ip), "routes": get_routes(ip),
"bridges": get_bridges(), # uses NDB internally "bridges": get_bridges(),
} }
@router.post("/bridge/create")
def create_bridge(req: BridgeCreateRequest, ip: IPRoute = Depends(get_iproute)):
if bridge_exists(req.name, ip):
raise HTTPException(400, detail=f"Bridge {req.name} already exists")
# Bridge erzeugen @router.post("/bridge/create")
def create_bridge(req: BridgeCreateRequest, ip: IPRoute = Depends(get_iproute)) -> dict:
"""Create a bridge and attach listed interfaces."""
if bridge_exists(req.name, ip):
raise HTTPException(status_code=400, detail=f"Bridge {req.name} already exists")
ip.link("add", ifname=req.name, kind="bridge") ip.link("add", ifname=req.name, kind="bridge")
br_idx = iface_index(req.name, ip) br_idx = iface_index(req.name, ip)
# Bridge konfigurieren
# TODO Parameter anpassen (STP, etc.)
ip.link("set", index=br_idx, kind="bridge", br_stp_state=0) ip.link("set", index=br_idx, kind="bridge", br_stp_state=0)
ip.link("set", index=br_idx, state="up") ip.link("set", index=br_idx, state="up")
# Interfaces hinzufügen + aktivieren
for iface in req.interfaces: for iface in req.interfaces:
idx = iface_index(iface, ip) idx = iface_index(iface, ip)
ip.link("set", index=idx, state="down")
# interface hochfahren
ip.link("set", index=idx, state="down") # optional - sicherer
ip.link("set", index=idx, state="up") ip.link("set", index=idx, state="up")
# interface in die bridge hängen
ip.link("set", index=idx, master=br_idx) ip.link("set", index=idx, master=br_idx)
return { return {
"status": "ok", "status": "ok",
"bridge": req.name, "bridge": req.name,
"interfaces": req.interfaces "interfaces": req.interfaces,
} }
@router.post("/bridge/remove") @router.post("/bridge/remove")
def remove_bridge(req: BridgeRemoveRequest, ip: IPRoute = Depends(get_iproute)): def remove_bridge(req: BridgeRemoveRequest, ip: IPRoute = Depends(get_iproute)) -> dict:
"""Detach and remove a bridge by name."""
if not bridge_exists(req.name, ip): if not bridge_exists(req.name, ip):
raise HTTPException(404, f"Bridge {req.name} not found") raise HTTPException(status_code=404, detail=f"Bridge {req.name} not found")
br_idx = iface_index(req.name, ip) br_idx = iface_index(req.name, ip)
# Bridge runterfahren
ip.link("set", index=br_idx, state="down") ip.link("set", index=br_idx, state="down")
# Bridge löschen
ip.link("del", index=br_idx) ip.link("del", index=br_idx)
return { return {
"status": "ok", "status": "ok",
"deleted": req.name "deleted": req.name,
} }

View File

@@ -1,96 +1,72 @@
# src/routers/packets.py """Packet history and streaming endpoints."""
import asyncio import asyncio
import base64 import base64
import json import json
import logging import logging
from typing import Optional, Any, Dict, List, Union from typing import Any, Dict, List, Optional, Union
from fastapi import APIRouter, HTTPException, Query, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
from pydantic import BaseModel from pydantic import BaseModel
from src.Models.packets import PacketDBModel
from fastapi import APIRouter, Query, WebSocket, WebSocketDisconnect, HTTPException
from fastapi.responses import JSONResponse
import src.shared_objects as shared import src.shared_objects as shared
from src.Models.packets import PacketDBModel
logger = logging.getLogger("packets_router") logger = logging.getLogger("packets_router")
router = APIRouter() router = APIRouter()
def _serialize_row_for_json(row: Union[Dict[str, Any], PacketDBModel, BaseModel]) -> Dict[str, Any]: def _serialize_row_for_json(row: Union[Dict[str, Any], PacketDBModel, BaseModel]) -> Dict[str, Any]:
""" """Convert one packet row to a JSON-safe dictionary."""
Convert a DB row or PacketDBModel into a JSON-serializable dict.
- If `row` is a Pydantic model (PacketDBModel or BaseModel), use `.dict()` to get a plain dict.
- If `raw` is bytes/bytearray, produce `raw_b64` and drop `raw`.
- If `raw_b64` already exists, do not re-encode.
- For values that cannot be JSON serialized, fall back to str(value).
"""
# If given a Pydantic model, convert to dict first
if isinstance(row, BaseModel): if isinstance(row, BaseModel):
d: Dict[str, Any] = row.dict(by_alias=True, exclude_none=True) raw_dict: Dict[str, Any] = row.dict(by_alias=True, exclude_none=True)
else: else:
# copy to avoid mutating caller's dict raw_dict = dict(row)
d = dict(row)
# If raw_b64 already present, prefer it. If raw present and bytes, convert. raw_val = raw_dict.get("raw")
raw_val = d.get("raw")
if raw_val is not None and isinstance(raw_val, (bytes, bytearray)): if raw_val is not None and isinstance(raw_val, (bytes, bytearray)):
# convert to base64 string and remove raw
try: try:
d["raw_b64"] = base64.b64encode(raw_val).decode("ascii") raw_dict["raw_b64"] = base64.b64encode(raw_val).decode("ascii")
d.pop("raw", None) raw_dict.pop("raw", None)
except Exception: except Exception:
# keep raw as str fallback
try: try:
d["raw_b64"] = base64.b64encode(bytes(raw_val)).decode("ascii") raw_dict["raw_b64"] = base64.b64encode(bytes(raw_val)).decode("ascii")
d.pop("raw", None) raw_dict.pop("raw", None)
except Exception: except Exception:
logger.exception("Failed to base64-encode raw bytes for row id=%s", d.get("id")) logger.exception("Failed to base64-encode raw bytes for row id=%s", raw_dict.get("id"))
d["raw_b64"] = str(raw_val) raw_dict["raw_b64"] = str(raw_val)
d.pop("raw", None) raw_dict.pop("raw", None)
# Ensure final dict is JSON-safe: try json.dumps on each value, fallback to str() output: Dict[str, Any] = {}
out: Dict[str, Any] = {} for key, value in raw_dict.items():
for k, v in d.items(): if key == "raw_b64" and isinstance(value, (bytes, bytearray)):
# skip any private/internal keys if needed (optional)
# if k.startswith("_"):
# continue
# raw_b64: ensure it's a str
if k == "raw_b64" and isinstance(v, (bytes, bytearray)):
try: try:
out["raw_b64"] = base64.b64encode(v).decode("ascii") output["raw_b64"] = base64.b64encode(value).decode("ascii")
continue
except Exception: except Exception:
out["raw_b64"] = str(v) output["raw_b64"] = str(value)
continue continue
# JSON-serializable check
try: try:
json.dumps({k: v}) json.dumps({key: value})
out[k] = v output[key] = value
except (TypeError, ValueError): except (TypeError, ValueError):
# convert non-serializable to string representation
try: try:
out[k] = str(v) output[key] = str(value)
except Exception: except Exception:
out[k] = "<unserializable>" output[key] = "<unserializable>"
return out
return output
async def _serialize_rows(rows: List[Union[Dict[str, Any], PacketDBModel]]) -> List[Dict[str, Any]]: async def _serialize_rows(rows: List[Union[Dict[str, Any], PacketDBModel]]) -> List[Dict[str, Any]]:
""" """Convert packet rows to JSON-safe dictionaries, preserving order."""
Serialize a list of DB rows or PacketDBModel instances into JSON-ready dicts. return [_serialize_row_for_json(row) for row in rows]
Keeps the same order as input.
"""
return [_serialize_row_for_json(r) for r in rows]
@router.get("/packets") @router.get("/packets")
async def get_packets(limit: int = Query(100, ge=1, le=10000)): async def get_packets(limit: int = Query(100, ge=1, le=10000)) -> JSONResponse:
""" """Return the latest packets in reverse chronological order."""
Return latest `limit` packets (newest first). The DB helper already converts
`raw` to `raw_b64` in fetch_latest, but we defensively re-serialize here.
"""
db = shared.db db = shared.db
if db is None: if db is None:
logger.warning("GET /packets called but DB is not available") logger.warning("GET /packets called but DB is not available")
@@ -98,21 +74,16 @@ async def get_packets(limit: int = Query(100, ge=1, le=10000)):
try: try:
rows = await db.fetch_latest(limit) rows = await db.fetch_latest(limit)
serial = await _serialize_rows(rows) serialized = await _serialize_rows(rows)
return JSONResponse(content={"count": len(serial), "packets": serial}) return JSONResponse(content={"count": len(serialized), "packets": serialized})
except Exception: except Exception as exc:
logger.exception("Failed to fetch latest packets from DB") logger.exception("Failed to fetch latest packets from DB")
raise HTTPException(status_code=500, detail="Failed to fetch packets") raise HTTPException(status_code=500, detail="Failed to fetch packets") from exc
@router.websocket("/ws/packets") @router.websocket("/ws/packets")
async def websocket_packets(ws: WebSocket): async def websocket_packets(ws: WebSocket) -> None:
""" """Stream live packets to a websocket client."""
WebSocket live feed endpoint.
Accepts optional query param `subscribe_recent` (e.g. ?subscribe_recent=20)
which will deliver the last N packets immediately on connect.
"""
await ws.accept() await ws.accept()
logger.debug("WebSocket connection accepted: %s", ws.client) logger.debug("WebSocket connection accepted: %s", ws.client)
@@ -131,86 +102,75 @@ async def websocket_packets(ws: WebSocket):
logger.warning("WebSocket closed: broadcaster not available") logger.warning("WebSocket closed: broadcaster not available")
return return
# Parse subscribe_recent from query params (defensive)
try: try:
subscribe_recent_raw = ws.query_params.get("subscribe_recent", "0") subscribe_recent_raw = ws.query_params.get("subscribe_recent", "0")
subscribe_recent = int(subscribe_recent_raw) subscribe_recent = int(subscribe_recent_raw)
if subscribe_recent < 0: subscribe_recent = max(subscribe_recent, 0)
subscribe_recent = 0
except Exception: except Exception:
subscribe_recent = 0 subscribe_recent = 0
q: Optional[asyncio.Queue] = None queue: Optional[asyncio.Queue] = None
try: try:
# Optionally send recent history first
if subscribe_recent > 0: if subscribe_recent > 0:
recent = await db.fetch_latest(subscribe_recent) recent = await db.fetch_latest(subscribe_recent)
recent_serial = await _serialize_rows(recent) recent_serialized = await _serialize_rows(recent)
await ws.send_json({"type": "recent", "count": len(recent_serial), "packets": recent_serial}) await ws.send_json({"type": "recent", "count": len(recent_serialized), "packets": recent_serialized})
# Subscribe to broadcaster to receive live packets queue = await broadcaster.subscribe()
q = await broadcaster.subscribe() logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, queue.maxsize)
logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, q.maxsize)
# Simple heartbeat: periodically ensure client is responsive (optional)
# We'll implement by awaiting q.get() which blocks until a message is published.
while True: while True:
msg = await q.get() message = await queue.get()
# Normalize message to JSON-able dict
if isinstance(msg, dict): if isinstance(message, dict):
payload = _serialize_row_for_json(msg) payload: Any = _serialize_row_for_json(message)
else: else:
# not a dict — try to json-serialize directly
try: try:
json.dumps(msg) json.dumps(message)
payload = msg payload = message
except Exception: except Exception:
payload = {"data": str(msg)} payload = {"data": str(message)}
try: try:
await ws.send_json(payload) await ws.send_json(payload)
except Exception: except Exception:
# sending failed (client disconnected or write error) logger.info("WebSocket send failed for client %s; unsubscribing", ws.client)
logger.info("WebSocket send failed for client %s — unsubscribing", ws.client)
break break
except WebSocketDisconnect: except WebSocketDisconnect:
logger.info("WebSocket client disconnected: %s", ws.client) logger.info("WebSocket client disconnected: %s", ws.client)
except Exception: except Exception:
logger.exception("Unexpected error in websocket_packets") logger.exception("Unexpected error in websocket_packets")
finally: finally:
# Clean up subscriber queue if queue is not None:
if q is not None:
try: try:
await broadcaster.unsubscribe(q) await broadcaster.unsubscribe(queue)
except Exception: except Exception:
logger.exception("Failed to unsubscribe websocket queue") logger.exception("Failed to unsubscribe websocket queue")
try: try:
await ws.close() await ws.close()
except Exception: except Exception:
pass pass
logger.debug("WebSocket connection closed and cleaned up for client %s", ws.client)
logger.debug("WebSocket connection cleaned up for client %s", ws.client)
@router.delete("/packets") @router.delete("/packets")
async def clear_packets(reset_id: bool = Query(True)): async def clear_packets(reset_id: bool = Query(True)) -> JSONResponse:
""" """Remove all packet rows from the database."""
Clear all packet logs from the database.
Uses TRUNCATE internally for high performance.
"""
db = shared.db db = shared.db
if db is None: if db is None:
raise HTTPException(status_code=503, detail="Database not available") raise HTTPException(status_code=503, detail="Database not available")
success = await db.clear_all_packets(reset_identity=reset_id) success = await db.clear_all_packets(reset_identity=reset_id)
if not success: if not success:
raise HTTPException(status_code=500, detail="Failed to clear packet table") raise HTTPException(status_code=500, detail="Failed to clear packet table")
logger.info("User initiated clear_packets (reset_id=%s)", reset_id) logger.info("User initiated clear_packets (reset_id=%s)", reset_id)
return JSONResponse( return JSONResponse(
content={ content={
"status": "success", "status": "success",
"message": "All packets have been cleared", "message": "All packets have been cleared",
"reset_id": reset_id "reset_id": reset_id,
} }
) )

View File

@@ -1,7 +1,9 @@
# src/routers/sniffer.py """HTTP API for starting, stopping, and inspecting sniffer sessions."""
from fastapi import APIRouter, HTTPException, Query, Body
from typing import Any, Dict, Optional
from fastapi import APIRouter, Body, HTTPException, Query
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from typing import Dict, Any, Optional
from src.network_sniffer import ( from src.network_sniffer import (
get_sniffer_status, get_sniffer_status,
@@ -12,103 +14,113 @@ from src.network_sniffer import (
router = APIRouter() router = APIRouter()
# ------------------------------
# Pydantic Models
# ------------------------------
class SnifferStartRequest(BaseModel): class SnifferStartRequest(BaseModel):
""" """Request payload for starting a sniffer session."""
Request model for starting the sniffer on a specific bridge OR interface.
Exactly one of `bridge` or `interface` must be provided.
"""
bridge: Optional[str] = Field(None, example="br0", description="Name of the Linux bridge to sniff on")
interface: Optional[str] = Field(None, example="eth0", description="Name of the network interface to sniff on")
bridge: Optional[str] = Field(
class SnifferStartResponse(BaseModel): None,
""" example="br0",
Response model returned when sniffer starts successfully. description="Bridge name to sniff.",
""" )
started: bool = Field(..., description="Whether the sniffer was started successfully") interface: Optional[str] = Field(
session_id: str = Field(..., description="Session identifier for this sniffer instance") None,
target: str = Field(..., description="Target that was started (bridge or interface)") example="eth0",
target_type: str = Field(..., description="Either 'bridge' or 'interface'") description="Interface name to sniff.",
)
class SnifferStopRequest(BaseModel):
""" class SnifferStartResponse(BaseModel):
Optional body for stop — prefer session_id if you want to stop a specific session. """Response payload for a successful sniffer start."""
If omitted, stopping behavior will be determined by query params (bridge/interface) or global stop.
""" started: bool = Field(..., description="True when a session was started.")
session_id: Optional[str] = Field(None, description="Session id to stop") session_id: str = Field(..., description="Unique session identifier.")
target: str = Field(..., description="Started target name.")
target_type: str = Field(..., description="Either 'bridge' or 'interface'.")
class SnifferStopResponse(BaseModel):
stopped: bool = Field(..., description="Whether the sniffer was stopped successfully")
session_id: Optional[str] = Field(None, description="Session id stopped (if any)") class SnifferStopRequest(BaseModel):
target: Optional[str] = Field(None, description="Target stopped; null if global stop") """Optional stop payload for targeting a specific session."""
target_type: Optional[str] = Field(None, description="'bridge' or 'interface' or None")
session_id: Optional[str] = Field(None, description="Session ID to stop.")
class InterfaceSnifferStatus(BaseModel):
running: bool = Field(..., description="Whether the sniffer thread/socket is active") class SnifferStopResponse(BaseModel):
exists: bool = Field(..., description="Whether the interface exists in /sys/class/net") """Response payload for stop operations."""
up: bool = Field(..., description="Whether the interface is operationally UP")
session_id: Optional[str] = Field(None, description="Session id owning this interface") stopped: bool = Field(..., description="True when stop completed.")
session_label: Optional[str] = Field(None, description="Human label for the session") session_id: Optional[str] = Field(None, description="Stopped session ID if available.")
target: Optional[str] = Field(None, description="Stopped target name.")
target_type: Optional[str] = Field(None, description="'bridge', 'interface', or null.")
class SnifferStatusResponse(BaseModel):
interfaces: Dict[str, InterfaceSnifferStatus] = Field(
..., description="Map of interface names to their sniffer status" class InterfaceSnifferStatus(BaseModel):
"""Status details for a single network interface."""
running: bool = Field(..., description="Whether a sniffer is currently active.")
exists: bool = Field(..., description="Whether the interface exists on the host.")
up: bool = Field(..., description="Whether the interface is operationally up.")
session_id: Optional[str] = Field(None, description="Owning sniffer session ID.")
session_label: Optional[str] = Field(None, description="Human-readable session label.")
class SnifferStatusResponse(BaseModel):
"""Status response keyed by interface name."""
interfaces: Dict[str, InterfaceSnifferStatus] = Field(
...,
description="Map of interface names to status objects.",
) )
# ------------------------------
# Endpoints
# ------------------------------
@router.post("/start", response_model=SnifferStartResponse) @router.post("/start", response_model=SnifferStartResponse)
def sniffer_start(req: SnifferStartRequest): def sniffer_start(req: SnifferStartRequest) -> SnifferStartResponse:
""" """Start one sniffer session for exactly one target."""
Start a sniffer session for the given bridge OR interface.
Exactly one of `bridge` or `interface` must be provided.
Returns a session_id to manage the session.
"""
if bool(req.bridge) == bool(req.interface): if bool(req.bridge) == bool(req.interface):
raise HTTPException(status_code=400, detail="Exactly one of 'bridge' or 'interface' must be provided") raise HTTPException(status_code=400, detail="Exactly one of 'bridge' or 'interface' must be provided")
try: try:
if req.interface: if req.interface:
session_id = start_afpacket_sniffer(req.interface, target_is_interface=True) session_id = start_afpacket_sniffer(req.interface, target_is_interface=True)
return SnifferStartResponse(started=True, session_id=session_id, target=req.interface, target_type="interface") return SnifferStartResponse(
else: started=True,
session_id = start_afpacket_sniffer(req.bridge, target_is_interface=False) session_id=session_id,
return SnifferStartResponse(started=True, session_id=session_id, target=req.bridge, target_type="bridge") target=req.interface,
target_type="interface",
)
session_id = start_afpacket_sniffer(req.bridge, target_is_interface=False)
return SnifferStartResponse(
started=True,
session_id=session_id,
target=req.bridge,
target_type="bridge",
)
except Exception as exc: except Exception as exc:
raise HTTPException(status_code=500, detail=f"Failed to start sniffer: {exc}") raise HTTPException(status_code=500, detail=f"Failed to start sniffer: {exc}") from exc
@router.post("/stop", response_model=SnifferStopResponse) @router.post("/stop", response_model=SnifferStopResponse)
def sniffer_stop( def sniffer_stop(
q_bridge: Optional[str] = Query(None, alias="bridge", description="If provided, stop sniffer sockets for this bridge"), q_bridge: Optional[str] = Query(
q_interface: Optional[str] = Query(None, alias="interface", description="If provided, stop sniffer socket for this interface"), None,
alias="bridge",
description="Stop sockets for this bridge.",
),
q_interface: Optional[str] = Query(
None,
alias="interface",
description="Stop sockets for this interface.",
),
body: SnifferStopRequest = Body(...), body: SnifferStopRequest = Body(...),
): ) -> SnifferStopResponse:
""" """Stop by session ID, target query, or globally when no selector is given."""
Stop sniffer sessions.
- If body.session_id is provided: stop that session (preferred).
- Else if query param `interface` provided: close socket for that interface across sessions.
- Else if query param `bridge` provided: remove snapshot / close sockets for that bridge across sessions.
- Else: stop all sessions (global stop).
"""
if body and body.session_id: if body and body.session_id:
try: try:
stop_afpacket_sniffer(session_id=body.session_id) stop_afpacket_sniffer(session_id=body.session_id)
return SnifferStopResponse(stopped=True, session_id=body.session_id, target=None, target_type=None) return SnifferStopResponse(stopped=True, session_id=body.session_id, target=None, target_type=None)
except Exception as exc: except Exception as exc:
raise HTTPException(status_code=500, detail=f"Failed to stop session {body.session_id}: {exc}") raise HTTPException(status_code=500, detail=f"Failed to stop session {body.session_id}: {exc}") from exc
# validate query params
if q_bridge and q_interface: if q_bridge and q_interface:
raise HTTPException(status_code=400, detail="Only one of 'bridge' or 'interface' may be provided") raise HTTPException(status_code=400, detail="Only one of 'bridge' or 'interface' may be provided")
@@ -116,29 +128,23 @@ def sniffer_stop(
if q_interface: if q_interface:
stop_afpacket_sniffer(target=q_interface, target_is_interface=True) stop_afpacket_sniffer(target=q_interface, target_is_interface=True)
return SnifferStopResponse(stopped=True, session_id=None, target=q_interface, target_type="interface") return SnifferStopResponse(stopped=True, session_id=None, target=q_interface, target_type="interface")
if q_bridge: if q_bridge:
stop_afpacket_sniffer(target=q_bridge, target_is_interface=False) stop_afpacket_sniffer(target=q_bridge, target_is_interface=False)
return SnifferStopResponse(stopped=True, session_id=None, target=q_bridge, target_type="bridge") return SnifferStopResponse(stopped=True, session_id=None, target=q_bridge, target_type="bridge")
# global stop
stop_afpacket_sniffer() stop_afpacket_sniffer()
return SnifferStopResponse(stopped=True, session_id=None, target=None, target_type=None) return SnifferStopResponse(stopped=True, session_id=None, target=None, target_type=None)
except Exception as exc: except Exception as exc:
raise HTTPException(status_code=500, detail=f"Failed to stop sniffer: {exc}") raise HTTPException(status_code=500, detail=f"Failed to stop sniffer: {exc}") from exc
@router.get("/status", response_model=SnifferStatusResponse) @router.get("/status", response_model=SnifferStatusResponse)
def sniffer_status(): def sniffer_status() -> SnifferStatusResponse:
""" """Return current sniffer status per interface."""
Return the sniffer status information.
"""
try: try:
raw = get_sniffer_status() raw: Dict[str, Dict[str, Any]] = get_sniffer_status()
# Convert raw dict → typed model typed = {key: InterfaceSnifferStatus(**value) for key, value in raw.items()}
typed = {
k: InterfaceSnifferStatus(**v)
for k, v in raw.items()
}
return SnifferStatusResponse(interfaces=typed) return SnifferStatusResponse(interfaces=typed)
except Exception as exc: except Exception as exc:
raise HTTPException(status_code=500, detail=f"Failed to query sniffer status: {exc}") raise HTTPException(status_code=500, detail=f"Failed to query sniffer status: {exc}") from exc

View File

@@ -1,30 +1,24 @@
# src/main.py """FastAPI application entrypoint and runtime wiring."""
import asyncio import asyncio
import logging import logging
import os 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_scripting_api
from src.api import nft_manager
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.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 import src.shared_objects as shared_objects
from src.api import nft_manager
from src.api import packet_api from src.api import packet_api
import src.api.nftables_api as nftables_api from src.api import packet_scripting_api
from src.utilities.database import DatabasePool
from src.utilities.packet_broadcaster import PacketBroadcaster
# ---- Config -----------------------------------------------------------
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
logging.basicConfig(level=logging.DEBUG) logging.basicConfig(level=logging.DEBUG)
# ---- Globals -----------------------------------------------
# Create DatabasePool instance (pool created on startup)
shared_objects.db = DatabasePool(DB_DSN) shared_objects.db = DatabasePool(DB_DSN)
app = FastAPI( app = FastAPI(
@@ -45,62 +39,44 @@ app.add_middleware(
allow_headers=["*"], allow_headers=["*"],
) )
# ---------------------
# Startup / Shutdown
# ---------------------
@app.on_event("startup") @app.on_event("startup")
async def on_startup(): async def on_startup() -> None:
""" """Initialize shared runtime objects on the FastAPI event loop."""
Initialize DB pool and broadcaster on the FastAPI event loop and
publish them into shared_objects so other modules (sniffer, routers)
can access them.
"""
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
shared_objects.web_loop = loop shared_objects.web_loop = loop
# Initialize DB pool bound to this loop
try: try:
await shared_objects.db.init_pool() await shared_objects.db.init_pool()
except Exception: except Exception:
logging.exception("Failed to initialize DB pool") logging.exception("Failed to initialize DB pool")
raise raise
# Create broadcaster and attach to DB so DB.insert_packet can publish updates
try: try:
shared_objects.broadcaster = PacketBroadcaster(loop) shared_objects.broadcaster = PacketBroadcaster(loop)
shared_objects.db.broadcaster = shared_objects.broadcaster shared_objects.db.broadcaster = shared_objects.broadcaster
except Exception: except Exception:
logging.exception("Failed to create/attach broadcaster") logging.exception("Failed to create/attach broadcaster")
# continue — DB is primary; broadcaster optional
# Drain any buffered packets from the sniffer (if it started earlier)
try: try:
# import sniffer here to avoid circular imports at module import time
from src import network_sniffer as sniffer from src import network_sniffer as sniffer
# sniffer provides drain_buffer_to_shared_db()
try: try:
sniffer.drain_buffer_to_shared_db() sniffer.drain_buffer_to_shared_db()
except Exception: except Exception:
logging.exception("Failed to drain sniffer buffer") logging.exception("Failed to drain sniffer buffer")
except ImportError: except ImportError:
# sniffer not present or not importable; skip logging.debug("Sniffer module not importable at startup; skipping buffer drain")
logging.debug("sniffer module not importable at startup; skipping buffer drain")
@app.on_event("shutdown") @app.on_event("shutdown")
async def shutdown_event(): async def shutdown_event() -> None:
""" """Stop network resources and release shared runtime objects."""
Shutdown actions: stop network API and close DB pool if present.
"""
# try to shut down network API components
try: try:
network_api.shutdown_network_api() network_api.shutdown_network_api()
except Exception: except Exception:
logging.exception("Error shutting down network API") logging.exception("Error shutting down network API")
# close DB pool if available in shared_objects
try: try:
web_db = getattr(shared_objects, "db", None) web_db = getattr(shared_objects, "db", None)
if web_db is not None: if web_db is not None:
@@ -108,37 +84,26 @@ async def shutdown_event():
except Exception: except Exception:
logging.exception("Failed to close DB pool during shutdown") logging.exception("Failed to close DB pool during shutdown")
# clear shared runtime objects (optional cleanup) shared_objects.db = None
try: shared_objects.broadcaster = None
shared_objects.db = None shared_objects.web_loop = None
shared_objects.broadcaster = None
shared_objects.web_loop = None
except Exception:
pass
# ---------------------
# Basic Endpoints
# ---------------------
@app.get("/hello") @app.get("/hello")
def hello(): def hello() -> dict[str, str]:
"""Simple health-check endpoint."""
return {"message": "Hello from FastAPI 🎉"} return {"message": "Hello from FastAPI 🎉"}
@app.get("/versions") @app.get("/versions")
def versions(): def versions() -> dict[str, str]:
"""Return runtime Python version."""
message = os.popen("python --version").read().strip() message = os.popen("python --version").read().strip()
return {"message": message} return {"message": message}
# ---------------------
# Routers
# ---------------------
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(nft_api.router, prefix="/nft", tags=["nft"])
app.include_router(nft_manager.router, tags=["firewall"]) app.include_router(nft_manager.router, tags=["firewall"])
app.include_router(packet_scripting_api.router, prefix="/scripts", tags=["scripts"]) app.include_router(packet_scripting_api.router, prefix="/scripts", tags=["scripts"])

View File

@@ -1,8 +1,8 @@
from typing import Optional """Shared runtime objects initialized during FastAPI startup."""
import asyncio
# These are filled at FastAPI startup import asyncio
# DB instance from typing import Any, Optional
db = None
db: Any = None
web_loop: Optional[asyncio.AbstractEventLoop] = None web_loop: Optional[asyncio.AbstractEventLoop] = None
broadcaster = None broadcaster: Any = None

View File

@@ -1,28 +1,21 @@
# src/utilities/database.py """Database helper for packet persistence and retrieval."""
import logging
import base64
import asyncio
from typing import Dict, List, Optional, Any
from pydantic import ValidationError
import asyncio
import base64
import logging
from typing import Any, Dict, List, Optional
import asyncpg import asyncpg
from asyncpg.pool import Pool from asyncpg.pool import Pool
from pydantic import ValidationError
from src.Models.packets import PacketDBModel from src.Models.packets import PacketDBModel
# ---- Logging ----------------------------------------------------------
logger = logging.getLogger("af_packet_sniffer") logger = logging.getLogger("af_packet_sniffer")
class DatabasePool: class DatabasePool:
""" """Asyncpg connection pool wrapper used by the packet APIs."""
Lightweight asyncpg connection pool wrapper.
- Lazy pool creation via init_pool()
- Safe against concurrent init_pool() calls via an asyncio.Lock created on first use
- insert_packet() forwards the pkt_info to an optional broadcaster after successful insert
"""
def __init__(self, dsn: str, min_size: int = 1, max_size: int = 5): def __init__(self, dsn: str, min_size: int = 1, max_size: int = 5):
self._dsn = dsn self._dsn = dsn
@@ -30,26 +23,26 @@ class DatabasePool:
self._min_size = min_size self._min_size = min_size
self._max_size = max_size self._max_size = max_size
self.broadcaster = None self.broadcaster = None
# created on first init_pool() (must be created on an event loop)
self._init_lock: Optional[asyncio.Lock] = None self._init_lock: Optional[asyncio.Lock] = None
async def init_pool(self) -> None: async def init_pool(self) -> None:
"""Initialize the asyncpg pool if not already initialized (idempotent).""" """Initialize the connection pool once per process."""
if self._pool is not None: if self._pool is not None:
return return
# Ensure a lock exists that is bound to the running event loop
if self._init_lock is None: if self._init_lock is None:
self._init_lock = asyncio.Lock() self._init_lock = asyncio.Lock()
async with self._init_lock: async with self._init_lock:
# Double-check after acquiring lock
if self._pool is not None: if self._pool is not None:
return return
logger.info("Initializing DB pool (dsn=%s)", self._dsn) logger.info("Initializing DB pool (dsn=%s)", self._dsn)
try: try:
self._pool = await asyncpg.create_pool( self._pool = await asyncpg.create_pool(
dsn=self._dsn, min_size=self._min_size, max_size=self._max_size dsn=self._dsn,
min_size=self._min_size,
max_size=self._max_size,
) )
logger.info("DB pool initialized") logger.info("DB pool initialized")
except Exception: except Exception:
@@ -57,9 +50,10 @@ class DatabasePool:
raise raise
async def close_pool(self) -> None: async def close_pool(self) -> None:
"""Close the pool if it exists.""" """Close the pool if present."""
if self._pool is None: if self._pool is None:
return return
try: try:
await self._pool.close() await self._pool.close()
logger.info("DB pool closed") logger.info("DB pool closed")
@@ -69,13 +63,7 @@ class DatabasePool:
self._pool = None self._pool = None
async def insert_packet(self, pkt_info: Dict[str, Any]) -> None: async def insert_packet(self, pkt_info: Dict[str, Any]) -> None:
""" """Insert one packet record and publish it to subscribers."""
Insert packet metadata into the `packets` table.
Preserves the same columns/values as before.
After a successful insert, if a broadcaster is attached it will be
notified via broadcaster.sync_publish(pkt_info).
"""
if self._pool is None: if self._pool is None:
await self.init_pool() await self.init_pool()
@@ -115,13 +103,11 @@ class DatabasePool:
except Exception: except Exception:
logger.exception("DB insert failed") logger.exception("DB insert failed")
return return
# Update the dictionary with the DB-generated values
if new_row: if new_row:
pkt_info["id"] = new_row["id"] pkt_info["id"] = new_row["id"]
# Convert timestamp to ISO string for JSON serialization in WebSockets
pkt_info["timestamp"] = new_row["timestamp"].isoformat() pkt_info["timestamp"] = new_row["timestamp"].isoformat()
# notify broadcaster (non-blocking). broadcaster is expected to be thread-safe.
if self.broadcaster: if self.broadcaster:
try: try:
self.broadcaster.sync_publish(pkt_info) self.broadcaster.sync_publish(pkt_info)
@@ -129,11 +115,7 @@ class DatabasePool:
logger.exception("Failed to publish pkt_info to broadcaster") logger.exception("Failed to publish pkt_info to broadcaster")
async def fetch_latest(self, limit: int) -> List[PacketDBModel]: async def fetch_latest(self, limit: int) -> List[PacketDBModel]:
""" """Fetch newest packet rows as validated `PacketDBModel` instances."""
Fetch the latest `limit` packets (newest first).
Returns a list of PacketDBModel. Converts raw bytes -> raw_b64 for JSON-safe output.
"""
if self._pool is None: if self._pool is None:
await self.init_pool() await self.init_pool()
@@ -148,45 +130,34 @@ class DatabasePool:
limit, limit,
) )
out: List[PacketDBModel] = [] result: List[PacketDBModel] = []
for row in rows:
data = dict(row)
for r in rows: raw_val = data.get("raw")
d = dict(r)
# convert byte raw -> base64 string (and remove raw)
raw_val = d.get("raw")
if isinstance(raw_val, (bytes, bytearray)): if isinstance(raw_val, (bytes, bytearray)):
d["raw_b64"] = base64.b64encode(raw_val).decode("ascii") data["raw_b64"] = base64.b64encode(raw_val).decode("ascii")
d.pop("raw", None) data.pop("raw", None)
# Validate/construct Pydantic model
try: try:
packet_model = PacketDBModel(**d) packet_model = PacketDBModel(**data)
except ValidationError as ve: except ValidationError as exc:
# Log and skip invalid rows (or handle otherwise)
logger.warning( logger.warning(
"Skipping DB row that failed PacketDBModel validation (id=%s): %s", "Skipping DB row that failed PacketDBModel validation (id=%s): %s",
d.get("id"), data.get("id"),
ve, exc,
) )
continue continue
out.append(packet_model) result.append(packet_model)
return result
return out
async def clear_all_packets(self, reset_identity: bool = True) -> bool: async def clear_all_packets(self, reset_identity: bool = True) -> bool:
""" """Truncate the packet table and optionally reset identity counters."""
Deletes all rows from the `packets` table.
If reset_identity is True, the auto-increment ID counter is reset to 1.
Returns True if successful, False otherwise.
"""
if self._pool is None: if self._pool is None:
await self.init_pool() await self.init_pool()
# TRUNCATE is faster than DELETE and resets the identity counter
restart_clause = "RESTART IDENTITY" if reset_identity else "" restart_clause = "RESTART IDENTITY" if reset_identity else ""
query = f"TRUNCATE TABLE packets {restart_clause};" query = f"TRUNCATE TABLE packets {restart_clause};"

View File

@@ -1,140 +1,100 @@
# src/utilities/packet_broadcaster.py """In-process packet broadcaster for websocket subscribers."""
import asyncio import asyncio
import logging import logging
from typing import Dict, Any, List, Optional from typing import Any, Dict, List, Optional
logger = logging.getLogger("packet_broadcaster") logger = logging.getLogger("packet_broadcaster")
class PacketBroadcaster: class PacketBroadcaster:
""" """Manage subscriber queues and publish packet events."""
Simple in-process broadcaster:
- Maintains a set of subscriber asyncio.Queues (one per websocket connection).
- publish(msg) is run on the broadcaster's event loop.
- sync_publish(msg) is thread-safe and can be called from other threads / loops.
Note: create this on the FastAPI event loop (e.g. in startup) so that its lock and
operations run on that same loop.
"""
def __init__(self, loop: asyncio.AbstractEventLoop, queue_maxsize: int = 1024): def __init__(self, loop: asyncio.AbstractEventLoop, queue_maxsize: int = 1024):
self._loop = loop self._loop = loop
self._queue_maxsize = queue_maxsize self._queue_maxsize = queue_maxsize
# create lock and subscribers on the target loop to avoid cross-loop asyncio primitives
self._subscribers: List[asyncio.Queue] = [] self._subscribers: List[asyncio.Queue] = []
# create lock bound to the same loop by scheduling its construction on that loop
self._lock: Optional[asyncio.Lock] = None self._lock: Optional[asyncio.Lock] = None
self._closed = False
try: try:
# ensure lock is created on the given loop def _make_lock() -> None:
def _make_lock():
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
loop.call_soon_threadsafe(_make_lock) loop.call_soon_threadsafe(_make_lock)
except Exception: except Exception:
# fallback — create in current loop if call_soon_threadsafe fails
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
self._closed = False
async def subscribe(self) -> asyncio.Queue: async def subscribe(self) -> asyncio.Queue:
""" """Create and register a queue for one subscriber."""
Create a subscriber queue and add it to the list.
Caller is expected to await on the returned queue to receive messages.
"""
if self._closed: if self._closed:
raise RuntimeError("PacketBroadcaster is closed") raise RuntimeError("PacketBroadcaster is closed")
q: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize) queue: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize)
# wait until lock exists
while self._lock is None: while self._lock is None:
await asyncio.sleep(0) # yield to event loop briefly await asyncio.sleep(0)
async with self._lock: async with self._lock:
self._subscribers.append(q) self._subscribers.append(queue)
return q
async def unsubscribe(self, q: asyncio.Queue) -> None: return queue
"""
Remove a subscriber queue if present. async def unsubscribe(self, queue: asyncio.Queue) -> None:
""" """Unregister a subscriber queue if it exists."""
if self._lock is None: if self._lock is None:
return return
async with self._lock: async with self._lock:
try: try:
self._subscribers.remove(q) self._subscribers.remove(queue)
except ValueError: except ValueError:
pass pass
async def publish(self, msg: Dict[str, Any]) -> None: async def publish(self, msg: Dict[str, Any]) -> None:
""" """Publish one message to all current subscribers."""
Publish msg to all subscribers (must be called on the broadcaster's loop). if self._closed or self._lock is None:
We use put_nowait to avoid blocking. If a subscriber queue is full we drop
that subscriber's message to avoid backpressure.
"""
if self._closed:
return
if self._lock is None:
# not initialized yet; nothing to do
return return
async with self._lock: async with self._lock:
subs = list(self._subscribers) subscribers = list(self._subscribers)
for q in subs: for queue in subscribers:
try: try:
q.put_nowait(msg) queue.put_nowait(msg)
except asyncio.QueueFull: except asyncio.QueueFull:
# drop message for this subscriber
continue continue
except Exception as exc: except Exception as exc:
logger.exception("Unexpected error when publishing to subscriber: %s", exc) logger.exception("Unexpected subscriber publish error: %s", exc)
# attempt to remove broken subscriber
try: try:
async with self._lock: async with self._lock:
if q in self._subscribers: if queue in self._subscribers:
self._subscribers.remove(q) self._subscribers.remove(queue)
except Exception: except Exception:
pass pass
def sync_publish(self, msg: Dict[str, Any]) -> None: def sync_publish(self, msg: Dict[str, Any]) -> None:
""" """Thread-safe wrapper that schedules `publish` on the broadcaster loop."""
Thread-safe publish method: schedule publish(msg) on the broadcaster's loop.
Safe to call from other threads / event loops.
We schedule creation of the publish task on the broadcaster loop using
call_soon_threadsafe so that publish() runs on the correct loop.
"""
if self._closed: if self._closed:
return return
try: try:
# schedule the coroutine to run on the broadcaster loop
self._loop.call_soon_threadsafe(asyncio.create_task, self.publish(msg)) self._loop.call_soon_threadsafe(asyncio.create_task, self.publish(msg))
except Exception as exc: except Exception as exc:
# swallow errors but log for debugging
logger.exception("sync_publish failed to schedule publish: %s", exc) logger.exception("sync_publish failed to schedule publish: %s", exc)
async def close(self) -> None: async def close(self) -> None:
""" """Close the broadcaster and clear queued messages."""
Close the broadcaster: mark closed, clear subscribers, and drain queues.
"""
self._closed = True self._closed = True
if self._lock is None: if self._lock is None:
return return
async with self._lock: async with self._lock:
subs = list(self._subscribers) subscribers = list(self._subscribers)
self._subscribers.clear() self._subscribers.clear()
for q in subs: for queue in subscribers:
try: try:
# optionally notify subscribers of closure by putting None (client must handle) while not queue.empty():
# q.put_nowait(None) queue.get_nowait()
while not q.empty():
try:
q.get_nowait()
except Exception:
break
except Exception: except Exception:
pass pass

View File

@@ -1,5 +1,5 @@
// src/apiClient.ts
import axios from 'axios'; import axios from 'axios';
import { CreateRuleRequest, ExecResult, RulesetModel } from '../types/firewall'; import { CreateRuleRequest, ExecResult, RulesetModel } from '../types/firewall';
import { import {
BridgeCreateRequest, BridgeCreateRequest,
@@ -36,22 +36,14 @@ export const api = axios.create({
timeout: 20000, timeout: 20000,
}); });
// Normalize FastAPI errors here
api.interceptors.response.use( api.interceptors.response.use(
(response) => response, (response) => response,
(error) => { (error) => {
// FastAPI HTTPException format
const detail = error?.response?.data?.detail ?? error?.response?.data?.message ?? error.message ?? 'Unknown error'; const detail = error?.response?.data?.detail ?? error?.response?.data?.message ?? error.message ?? 'Unknown error';
// Always reject with a standard Error
return Promise.reject(new Error(detail)); return Promise.reject(new Error(detail));
}, },
); );
/* -------------------------
Basic endpoints
------------------------- */
export const fetchHello = async (): Promise<any> => { export const fetchHello = async (): Promise<any> => {
const res = await api.get('/hello'); const res = await api.get('/hello');
return res.data; return res.data;
@@ -62,10 +54,6 @@ export const fetchVersions = async (): Promise<any> => {
return res.data; return res.data;
}; };
/* -------------------------
Network queries (existing)
------------------------- */
export const fetchInterfaces = async (): Promise<InterfaceInfo[]> => { export const fetchInterfaces = async (): Promise<InterfaceInfo[]> => {
const res = await api.get<InterfaceInfo[]>('/network/interfaces'); const res = await api.get<InterfaceInfo[]>('/network/interfaces');
return res.data; return res.data;
@@ -101,150 +89,84 @@ export const removeBridge = async (req: BridgeRemoveRequest) => {
return res.data; return res.data;
}; };
/* -------------------------
Sniffer
------------------------- */
/**
* Start a sniffer session. Provide exactly one of { bridge, interface }.
* Returns a session_id that you can use to stop the session later.
*/
export const startSniffer = async (payload: SnifferStartRequest): Promise<SnifferStartResponse> => { export const startSniffer = async (payload: SnifferStartRequest): Promise<SnifferStartResponse> => {
const res = await api.post<SnifferStartResponse>('/sniffer/start', payload); const res = await api.post<SnifferStartResponse>('/sniffer/start', payload);
return res.data; return res.data;
}; };
/**
* Stop sniffer(s).
* - If you pass a body with { session_id }, it will stop that specific session.
* - If you call without body and without query params, it will stop all sessions.
*/
export const stopSniffer = async (body?: SnifferStopRequest): Promise<SnifferStopResponse> => { export const stopSniffer = async (body?: SnifferStopRequest): Promise<SnifferStopResponse> => {
const res = await api.post<SnifferStopResponse>('/sniffer/stop', body ?? {}); const res = await api.post<SnifferStopResponse>('/sniffer/stop', body ?? {});
return res.data; return res.data;
}; };
/**
* Stop sniffing for a specific interface across sessions.
* Calls: POST /sniffer/stop?interface=eth0 (empty body)
*/
export const stopSnifferByInterface = async (iface: string): Promise<SnifferStopResponse> => { export const stopSnifferByInterface = async (iface: string): Promise<SnifferStopResponse> => {
const res = await api.post<SnifferStopResponse>(`/sniffer/stop?interface=${encodeURIComponent(iface)}`, {}); const res = await api.post<SnifferStopResponse>(`/sniffer/stop?interface=${encodeURIComponent(iface)}`, {});
return res.data; return res.data;
}; };
/**
* Stop sniffing for a specific bridge across sessions.
* Calls: POST /sniffer/stop?bridge=br0 (empty body)
*/
export const stopSnifferByBridge = async (bridge: string): Promise<SnifferStopResponse> => { export const stopSnifferByBridge = async (bridge: string): Promise<SnifferStopResponse> => {
const res = await api.post<SnifferStopResponse>(`/sniffer/stop?bridge=${encodeURIComponent(bridge)}`, {}); const res = await api.post<SnifferStopResponse>(`/sniffer/stop?bridge=${encodeURIComponent(bridge)}`, {});
return res.data; return res.data;
}; };
/**
* Fetch the sniffer status (per-interface).
*/
export const fetchSnifferStatus = async (): Promise<SnifferStatusResponse> => { export const fetchSnifferStatus = async (): Promise<SnifferStatusResponse> => {
const res = await api.get<SnifferStatusResponse>('/sniffer/status'); const res = await api.get<SnifferStatusResponse>('/sniffer/status');
return res.data; return res.data;
}; };
/* -------------------------
Packets
------------------------- */
export const fetchPackets = async (limit = 100): Promise<any> => { export const fetchPackets = async (limit = 100): Promise<any> => {
// limit default mirrors OpenAPI default
const res = await api.get('/packets/packets', { params: { limit } }); const res = await api.get('/packets/packets', { params: { limit } });
return res.data; return res.data;
}; };
export const clearPackets = async (): Promise<any> => { export const clearPackets = async (): Promise<any> => {
// limit default mirrors OpenAPI default
const res = await api.delete('/packets/packets'); const res = await api.delete('/packets/packets');
return res.data; return res.data;
}; };
/* -------------------------
Firewall
------------------------- */
/**
* GET /firewall/rules
* Returns: { ruleset: RulesetModel | string | null }
* - If the server returns a raw textual fallback (string), the caller should handle it.
*/
export const fetchRuleset = async (): Promise<{ ruleset: RulesetModel }> => { export const fetchRuleset = async (): Promise<{ ruleset: RulesetModel }> => {
const res = await api.get<{ ruleset: RulesetModel }>('/firewall/rules'); const res = await api.get<{ ruleset: RulesetModel }>('/firewall/rules');
return res.data; return res.data;
}; };
/**
* DELETE /firewall/rules/{handle}?family=...&table=...&chain=...
* On success the backend returns 204 No Content. This function resolves to void.
*/
export const deleteRule = async (handle: number, family: string, table: string, chain: string): Promise<void> => { export const deleteRule = async (handle: number, family: string, table: string, chain: string): Promise<void> => {
const res = await api.delete(`/firewall/rules/${encodeURIComponent(String(handle))}`, { const res = await api.delete(`/firewall/rules/${encodeURIComponent(String(handle))}`, {
params: { family, table, chain }, params: { family, table, chain },
}); });
// axios resolves non-2xx as reject; server uses 204 No Content so nothing to return
return res.data; return res.data;
}; };
/**
* createRuleJson - POST /firewall/rules
* Body: CreateRuleRequest (must include expr)
*/
export const createRuleJson = async (req: CreateRuleRequest): Promise<ExecResult> => { export const createRuleJson = async (req: CreateRuleRequest): Promise<ExecResult> => {
const res = await api.post<ExecResult>('/firewall/rules', req); const res = await api.post<ExecResult>('/firewall/rules', req);
return res.data; return res.data;
}; };
export const execFirewallRaw = async (cmd: string): Promise<ExecResult> => { export const execFirewallRaw = async (cmd: string): Promise<ExecResult> => {
const res = await api.post<ExecResult>('/firewall/raw', { cmd: cmd }); const res = await api.post<ExecResult>('/firewall/raw', { cmd });
return res.data; return res.data;
}; };
/* -------------------------
Scripts
------------------------- */
/**
* Fetch the combined scripts + status endpoint.
* Returns a list of ScriptWithStatus entries.
*/
export const fetchScriptsAll = async (): Promise<ScriptWithStatus[]> => { export const fetchScriptsAll = async (): Promise<ScriptWithStatus[]> => {
const res = await api.get<ScriptWithStatus[]>('/scripts/scripts'); const res = await api.get<ScriptWithStatus[]>('/scripts/scripts');
return res.data; return res.data;
}; };
/**
* Convenience: list only the basic ScriptInfo items (no mappings).
* This uses the combined endpoint and maps to ScriptInfo[].
*/
export const listScripts = async (): Promise<ScriptInfo[]> => { export const listScripts = async (): Promise<ScriptInfo[]> => {
const all = await fetchScriptsAll(); const all = await fetchScriptsAll();
return all.map((s) => ({ name: s.name, path: s.path })); return all.map((script) => ({ name: script.name, path: script.path }));
}; };
/**
* Get status/mappings for a single script by name.
* Because the backend merged status into /scripts, we fetch that and filter.
*/
export const fetchScriptStatusForName = async (name: string): Promise<StatusForNameResponse> => { export const fetchScriptStatusForName = async (name: string): Promise<StatusForNameResponse> => {
const all = await fetchScriptsAll(); const all = await fetchScriptsAll();
const found = all.find((s) => s.name === name); const found = all.find((script) => script.name === name);
if (!found) { if (!found) {
// If script not found we still return an empty mapping structure
return { name, mappings: [] }; return { name, mappings: [] };
} }
return { name: found.name, mappings: found.mappings || [] }; return { name: found.name, mappings: found.mappings || [] };
}; };
/**
* Upload a script (multipart). If a requirements file is provided the backend
* will run pip install and return pip output in the response (mandatory install).
*/
export const uploadScript = async (opts: { export const uploadScript = async (opts: {
name: string; name: string;
script: File | Blob; script: File | Blob;
@@ -253,7 +175,10 @@ export const uploadScript = async (opts: {
const fd = new FormData(); const fd = new FormData();
fd.append('name', opts.name); fd.append('name', opts.name);
fd.append('script', opts.script); fd.append('script', opts.script);
if (opts.requirements) fd.append('requirements', opts.requirements as Blob);
if (opts.requirements) {
fd.append('requirements', opts.requirements as Blob);
}
const res = await api.post<ScriptUploadResponse>('/scripts/scripts', fd, { const res = await api.post<ScriptUploadResponse>('/scripts/scripts', fd, {
headers: { 'Content-Type': 'multipart/form-data' }, headers: { 'Content-Type': 'multipart/form-data' },
@@ -261,71 +186,49 @@ export const uploadScript = async (opts: {
return res.data; return res.data;
}; };
/**
* Download the script file as a Blob. Use this blob to create an object URL or read text.
*/
export const downloadScript = async (name: string): Promise<Blob> => { export const downloadScript = async (name: string): Promise<Blob> => {
const res = await api.get(`/scripts/scripts/${encodeURIComponent(name)}`, { responseType: 'blob' }); const res = await api.get(`/scripts/scripts/${encodeURIComponent(name)}`, { responseType: 'blob' });
return res.data as Blob; return res.data as Blob;
}; };
/**
* Download the requirements.txt for a script as a Blob.
* Returns 404 if not present (axios will throw).
*/
export const downloadRequirements = async (name: string): Promise<Blob> => { export const downloadRequirements = async (name: string): Promise<Blob> => {
const res = await api.get(`/scripts/scripts/${encodeURIComponent(name)}/requirements`, { responseType: 'blob' }); const res = await api.get(`/scripts/scripts/${encodeURIComponent(name)}/requirements`, { responseType: 'blob' });
return res.data as Blob; return res.data as Blob;
}; };
/**
* Replace / upload requirements for a script. This endpoint ALWAYS runs pip install
* and will return pip stdout/stderr on success. On pip failure the backend returns 500.
*
* Use multipart/form-data so Swagger UI shows a file picker (backend expects UploadFile).
*/
export const uploadRequirements = async ( export const uploadRequirements = async (
name: string, name: string,
requirements: File | Blob, requirements: File | Blob,
): Promise<RequirementsUploadResult> => { ): Promise<RequirementsUploadResult> => {
const fd = new FormData(); const fd = new FormData();
fd.append('requirements', requirements); fd.append('requirements', requirements);
const res = await api.put<RequirementsUploadResult>(`/scripts/scripts/${encodeURIComponent(name)}/requirements`, fd, { const res = await api.put<RequirementsUploadResult>(`/scripts/scripts/${encodeURIComponent(name)}/requirements`, fd, {
headers: { 'Content-Type': 'multipart/form-data' }, headers: { 'Content-Type': 'multipart/form-data' },
}); });
return res.data; return res.data;
}; };
/**
* Delete only the requirements file and attempt to clean up the venv.
* Returns summary `{ removed: { requirements_removed, venv_removed }, errors? }`.
*/
export const deleteRequirements = async (name: string): Promise<RequirementsDeleteResult> => { export const deleteRequirements = async (name: string): Promise<RequirementsDeleteResult> => {
const res = await api.delete<RequirementsDeleteResult>(`/scripts/scripts/${encodeURIComponent(name)}/requirements`); const res = await api.delete<RequirementsDeleteResult>(`/scripts/scripts/${encodeURIComponent(name)}/requirements`);
return res.data; return res.data;
}; };
/**
* Delete a script (and its files / venv). If qnum provided only remove that unit.
*/
export const deleteScript = async (name: string, qnum?: number | null): Promise<DeleteResult> => { export const deleteScript = async (name: string, qnum?: number | null): Promise<DeleteResult> => {
const params: Record<string, any> = {}; const params: Record<string, any> = {};
if (typeof qnum !== 'undefined' && qnum !== null) params.qnum = qnum; if (typeof qnum !== 'undefined' && qnum !== null) {
params.qnum = qnum;
}
const res = await api.delete<DeleteResult>(`/scripts/scripts/${encodeURIComponent(name)}`, { params }); const res = await api.delete<DeleteResult>(`/scripts/scripts/${encodeURIComponent(name)}`, { params });
return res.data; return res.data;
}; };
/**
* Enable a script (creates + starts a systemd unit). Returns OperationResult including service name.
*/
export const enableScript = async (name: string, req: EnableRequest): Promise<OperationResult> => { export const enableScript = async (name: string, req: EnableRequest): Promise<OperationResult> => {
const res = await api.post<OperationResult>(`/scripts/scripts/${encodeURIComponent(name)}/enable`, req); const res = await api.post<OperationResult>(`/scripts/scripts/${encodeURIComponent(name)}/enable`, req);
return res.data; return res.data;
}; };
/**
* Disable a script unit (POST with qnum as query param). Returns OperationResult.
*/
export const disableScript = async (name: string, qnum: number): Promise<OperationResult> => { export const disableScript = async (name: string, qnum: number): Promise<OperationResult> => {
const res = await api.post<OperationResult>(`/scripts/scripts/${encodeURIComponent(name)}/disable`, null, { const res = await api.post<OperationResult>(`/scripts/scripts/${encodeURIComponent(name)}/disable`, null, {
params: { qnum }, params: { qnum },

View File

@@ -1,6 +1,6 @@
// src/AppRouter.tsx
import { Navigate, Route, Routes } from 'react-router-dom'; import { Navigate, Route, Routes } from 'react-router-dom';
import App from './App'; // your layout component (has <Outlet />)
import App from './App';
import { Firewall } from './pages/Firewall'; import { Firewall } from './pages/Firewall';
import Home from './pages/Home'; import Home from './pages/Home';
import Network from './pages/Network'; import Network from './pages/Network';
@@ -8,34 +8,27 @@ import Scripting from './pages/Scripting';
import Sniffing from './pages/Sniffing'; import Sniffing from './pages/Sniffing';
import { PATHS } from './routes'; import { PATHS } from './routes';
function NotFound() {
return (
<div style={{ padding: 16 }}>
<h2>404 - Not Found</h2>
<p>The requested page does not exist.</p>
</div>
);
}
export default function AppRouter() { export default function AppRouter() {
return ( return (
<Routes> <Routes>
{/* App is the top-level layout; Outlet renders the active child route */}
<Route path={PATHS.ROOT} element={<App />}> <Route path={PATHS.ROOT} element={<App />}>
{/* When the user hits '/', redirect to '/home' */}
<Route index element={<Navigate to={PATHS.HOME} replace />} /> <Route index element={<Navigate to={PATHS.HOME} replace />} />
{/* Child routes - these render inside App's <Outlet /> */}
<Route path={PATHS.HOME.slice(1)} element={<Home />} /> <Route path={PATHS.HOME.slice(1)} element={<Home />} />
<Route path={PATHS.NETWORK.slice(1)} element={<Network />} /> <Route path={PATHS.NETWORK.slice(1)} element={<Network />} />
<Route path={PATHS.SNIFFING.slice(1)} element={<Sniffing />} /> <Route path={PATHS.SNIFFING.slice(1)} element={<Sniffing />} />
<Route path={PATHS.SCRIPTING.slice(1)} element={<Scripting />} /> <Route path={PATHS.SCRIPTING.slice(1)} element={<Scripting />} />
<Route path={PATHS.FIREWALL.slice(1)} element={<Firewall />} /> <Route path={PATHS.FIREWALL.slice(1)} element={<Firewall />} />
{/* Fallback (renders inside layout too) */}
<Route path="*" element={<NotFound />} /> <Route path="*" element={<NotFound />} />
</Route> </Route>
</Routes> </Routes>
); );
} }
/** simple 404 rendered inside the layout */
function NotFound() {
return (
<div style={{ padding: 16 }}>
<h2>404 – Not Found</h2>
<p>The requested page does not exist.</p>
</div>
);
}

View File

@@ -1,4 +1,3 @@
// src/components/AddChainModal.tsx
import { CopyOutlined } from '@ant-design/icons'; import { CopyOutlined } from '@ant-design/icons';
import { import {
Alert, Alert,
@@ -67,12 +66,10 @@ export default function FirewallAddChainModal({
const isPrefilled = Boolean(table?.family && table?.name); const isPrefilled = Boolean(table?.family && table?.name);
// init form values when modal opens or table prop changes
useEffect(() => { useEffect(() => {
form.setFieldsValue({ form.setFieldsValue({
family: table?.family ?? 'bridge', family: table?.family ?? 'bridge',
tableName: table?.name ?? 'filter', tableName: table?.name ?? 'filter',
// Do not override chainName if the user has typed it previously
type: 'filter', type: 'filter',
hook: 'forward', hook: 'forward',
priority: 0, priority: 0,
@@ -111,7 +108,6 @@ export default function FirewallAddChainModal({
} }
} }
// Build add chain command — respects provided values or live form values
function buildCommands(values?: any): string[] { function buildCommands(values?: any): string[] {
const vals = values ?? form.getFieldsValue(); const vals = values ?? form.getFieldsValue();
@@ -128,7 +124,6 @@ export default function FirewallAddChainModal({
: 0; : 0;
const policy = vals.policy ?? ''; const policy = vals.policy ?? '';
// chain name: prefer explicit chainName else fallback to hook (less ideal) else 'mychain'
const chain = hook; const chain = hook;
const policyPart = policy ? ` policy ${policy} ;` : ''; const policyPart = policy ? ` policy ${policy} ;` : '';
@@ -136,7 +131,6 @@ export default function FirewallAddChainModal({
return [cmd]; return [cmd];
} }
// Execute commands sequentially
async function executeCommands(cmds: string[]) { async function executeCommands(cmds: string[]) {
setRunning(true); setRunning(true);
setResults([]); setResults([]);
@@ -156,24 +150,19 @@ export default function FirewallAddChainModal({
setResults(acc); setResults(acc);
setRunning(false); setRunning(false);
// refresh ruleset shown in modal
try { try {
await refreshRuleset(); await refreshRuleset();
} catch { } catch {
/* ignored, refreshRuleset already notified on failure */
} }
const hadError = acc.some((r) => r.err); const hadError = acc.some((r) => r.err);
if (!hadError) { if (!hadError) {
// success notification inside the modal (not global refresh notification)
notification.success({ notification.success({
message: 'Chain created', message: 'Chain created',
description: 'Chain created and ruleset refreshed locally in the modal.', description: 'Chain created and ruleset refreshed locally in the modal.',
duration: 4, duration: 4,
}); });
// Inform parent that a resource was created.
// Parent can decide whether to refresh and whether to show a notification.
onClose?.(true); onClose?.(true);
if (onSuccess) onSuccess(); if (onSuccess) onSuccess();
} else { } else {
@@ -219,13 +208,11 @@ export default function FirewallAddChainModal({
type="primary" type="primary"
onClick={async () => { onClick={async () => {
try { try {
// validate name fields: chainName required; if not prefilled, tableName & family required
const requiredFields = ['chainName']; const requiredFields = ['chainName'];
if (!isPrefilled) requiredFields.push('tableName', 'family'); if (!isPrefilled) requiredFields.push('tableName', 'family');
await form.validateFields(requiredFields as any); await form.validateFields(requiredFields as any);
setStep(1); setStep(1);
} catch { } catch {
// AntD will show validation messages; no extra handling required
} }
}} }}
> >
@@ -264,7 +251,7 @@ export default function FirewallAddChainModal({
footer={renderFooter()} footer={renderFooter()}
destroyOnClose destroyOnClose
> >
{/* Step 0: Form */}
{step === 0 && ( {step === 0 && (
<div> <div>
{rulesetEmpty && ( {rulesetEmpty && (
@@ -284,7 +271,7 @@ export default function FirewallAddChainModal({
}} }}
> >
<Row gutter={12}> <Row gutter={12}>
{/* family & table: show inputs only when not provided via props */}
<Col span={8}> <Col span={8}>
{isPrefilled ? ( {isPrefilled ? (
<Form.Item label="Family"> <Form.Item label="Family">
@@ -363,7 +350,7 @@ export default function FirewallAddChainModal({
</div> </div>
)} )}
{/* Step 1: Preview & Results */}
{step === 1 && ( {step === 1 && (
<> <>
<Card title="Command preview" style={{ marginBottom: 12 }}> <Card title="Command preview" style={{ marginBottom: 12 }}>

View File

@@ -1,4 +1,3 @@
// src/components/FirewallManager.tsx
import { Button, Card, Col, Form, Input, Modal, Row, Select, Space, Spin, Typography, notification } from 'antd'; import { Button, Card, Col, Form, Input, Modal, Row, Select, Space, Spin, Typography, notification } from 'antd';
import { ReactElement, useEffect, useState } from 'react'; import { ReactElement, useEffect, useState } from 'react';
import { execFirewallRaw, fetchRuleset } from '../api/apiClient'; import { execFirewallRaw, fetchRuleset } from '../api/apiClient';
@@ -42,7 +41,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
const [results, setResults] = useState<CmdResult[]>([]); const [results, setResults] = useState<CmdResult[]>([]);
const [selectedFamily, setSelectedFamily] = useState<string>('bridge'); const [selectedFamily, setSelectedFamily] = useState<string>('bridge');
// initialize and refresh when modal opens
useEffect(() => { useEffect(() => {
form.setFieldsValue({ form.setFieldsValue({
family: 'bridge', family: 'bridge',
@@ -70,7 +68,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
setRulesetEmpty(false); setRulesetEmpty(false);
} }
} catch (err: any) { } catch (err: any) {
// use notification instead of message
notification.warning({ notification.warning({
message: 'Could not load ruleset', message: 'Could not load ruleset',
description: err?.message ?? String(err), description: err?.message ?? String(err),
@@ -82,14 +79,12 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
} }
} }
// Build commands: only create table
function buildCommands(values: any): string[] { function buildCommands(values: any): string[] {
const family = (values.family ?? 'bridge').trim(); const family = (values.family ?? 'bridge').trim();
const table = (values.tableName ?? 'filter').trim(); const table = (values.tableName ?? 'filter').trim();
return [`add table ${family} ${table}`]; return [`add table ${family} ${table}`];
} }
// Execute commands sequentially
async function executeCommands(cmds: string[]) { async function executeCommands(cmds: string[]) {
setRunning(true); setRunning(true);
setResults([]); setResults([]);
@@ -109,16 +104,13 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
setResults(acc); setResults(acc);
setRunning(false); setRunning(false);
// refresh ruleset after running
try { try {
await refreshRuleset(); await refreshRuleset();
} catch { } catch {
// ignore — refreshRuleset handles notifications on error
} }
const hadError = acc.some((r) => r.err); const hadError = acc.some((r) => r.err);
if (!hadError) { if (!hadError) {
// success -> notify briefly and inform parent that a resource was created
notification.success({ notification.success({
message: 'Table created', message: 'Table created',
description: 'Table was created and ruleset has been refreshed locally in the modal.', description: 'Table was created and ruleset has been refreshed locally in the modal.',
@@ -126,7 +118,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
}); });
onClose?.(true); // signal parent to refresh and close modal onClose?.(true); // signal parent to refresh and close modal
} else { } else {
// error -> open results view and show notification
notification.error({ notification.error({
message: 'Some commands returned errors', message: 'Some commands returned errors',
description: 'See execution results below for details.', description: 'See execution results below for details.',
@@ -153,7 +144,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
<Button <Button
type="primary" type="primary"
onClick={() => { onClick={() => {
// validate form before preview
form form
.validateFields() .validateFields()
.then(() => setStep(1)) .then(() => setStep(1))
@@ -181,7 +171,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
); );
} }
// show spinner while checking ruleset
if (localLoadingRuleset) { if (localLoadingRuleset) {
return ( return (
<Modal title="Create Table" open={open} onCancel={() => onClose?.(false)} footer={null} width={700}> <Modal title="Create Table" open={open} onCancel={() => onClose?.(false)} footer={null} width={700}>
@@ -201,7 +190,7 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
footer={renderFooter()} footer={renderFooter()}
destroyOnClose destroyOnClose
> >
{/* Step 0: minimal form */}
{step === 0 && ( {step === 0 && (
<div> <div>
<Form <Form
@@ -238,7 +227,7 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview =
</div> </div>
)} )}
{/* Step 1: preview + results */}
{step === 1 && ( {step === 1 && (
<> <>
<Card title="Command preview" style={{ marginBottom: 12 }}> <Card title="Command preview" style={{ marginBottom: 12 }}>

View File

@@ -1,14 +1,3 @@
// src/components/RuleBuilder.tsx
//
// Extended RuleBuilder using the user's canonical match list:
// 1) Metadata & Connection Tracking (meta, ct)
// 2) Layer 3 Network Headers (ip, ip6)
// 3) Layer 4 Transport Headers (tcp, udp, icmp) — appear when chosen
// 4) Layer 2 Ethernet & VLAN (ether, vlan)
//
// The UI provides rich dropdowns / placeholders / short explanations for every token subfield.
//
// NOTE: This file replaces the token lists and per-field UI to strictly follow the user's canonical list.
import { PlusOutlined } from '@ant-design/icons'; import { PlusOutlined } from '@ant-design/icons';
import { import {
@@ -38,9 +27,6 @@ const { Title, Text } = Typography;
type FormValues = Record<string, any>; type FormValues = Record<string, any>;
/* ----------------------
Token types (canonical per user)
---------------------- */
type TokenType = type TokenType =
| 'meta' | 'meta'
| 'ct' | 'ct'
@@ -59,642 +45,10 @@ type TokenType =
| 'nat' | 'nat'
| 'queue'; | 'queue';
/* ----------------------
TOKEN_FIELD_OPTIONS
Each token lists allowed subfields (exactly the fields from the user's canonical list).
The `kind` tells the UI which input widget to show (number, enum, string).
---------------------- */
const TOKEN_FIELD_OPTIONS: Record< const TOKEN_FIELD_OPTIONS: Record<
TokenType, TokenType,
Array<{ value: string; label: string; kind?: 'string' | 'number' | 'enum' }> Array<{ value: string; label: string; kind?: 'string' | 'number' | 'enum' }>
> = { > =
/* 1) Metadata & Connection Tracking */
meta: [
{ value: 'iifname', label: 'iifname (input interface)', kind: 'string' },
{ value: 'oifname', label: 'oifname (output interface)', kind: 'string' },
{ value: 'l4proto', label: 'l4proto (protocol L4)', kind: 'enum' }, // tcp/udp/icmp/...
{ value: 'day', label: 'day (day of week)', kind: 'enum' },
{ value: 'hour', label: 'hour (hour of day/range)', kind: 'string' },
{ value: 'pkttype', label: 'pkttype (packet type)', kind: 'enum' },
{ value: 'mark', label: 'mark (packet mark)', kind: 'string' },
{ value: 'skuid', label: 'skuid (socket UID)', kind: 'number' },
{ value: 'skgid', label: 'skgid (socket GID)', kind: 'number' },
],
ct: [
{ value: 'state', label: 'state (ct state)', kind: 'enum' },
{ value: 'direction', label: 'direction (original/reply)', kind: 'enum' },
{ value: 'status', label: 'status', kind: 'string' },
{ value: 'mark', label: 'mark (conntrack mark)', kind: 'string' },
{ value: 'expiration', label: 'expiration', kind: 'string' },
{ value: 'helper', label: 'helper', kind: 'string' },
],
/* 2) Layer 3: Network Headers */
ip: [
{ value: 'saddr', label: 'saddr (source IPv4)', kind: 'string' },
{ value: 'daddr', label: 'daddr (destination IPv4)', kind: 'string' },
{ value: 'protocol', label: 'protocol (L4) — alias to l4proto', kind: 'enum' },
{ value: 'dscp', label: 'dscp (DSCP)', kind: 'enum' },
{ value: 'ttl', label: 'ttl (time to live)', kind: 'number' },
{ value: 'frag-off', label: 'frag-off (fragment bits)', kind: 'string' },
],
ip6: [
{ value: 'saddr', label: 'saddr (source IPv6)', kind: 'string' },
{ value: 'daddr', label: 'daddr (destination IPv6)', kind: 'string' },
{ value: 'nexthdr', label: 'nexthdr (protocol / next header)', kind: 'enum' },
{ value: 'dscp', label: 'dscp (DSCP)', kind: 'enum' },
{ value: 'hoplimit', label: 'hoplimit (IPv6 hop limit)', kind: 'number' },
{ value: 'flowlabel', label: 'flowlabel', kind: 'number' },
],
/* 3) Layer 4: Transport Headers (appear only when token type tcp/udp/icmp is chosen) */
tcp: [
{ value: 'sport', label: 'sport (source port)', kind: 'number' },
{ value: 'dport', label: 'dport (destination port)', kind: 'number' },
{ value: 'flags', label: 'flags (tcp flags bitmask)', kind: 'enum' },
],
udp: [
{ value: 'sport', label: 'sport (source port)', kind: 'number' },
{ value: 'dport', label: 'dport (destination port)', kind: 'number' },
],
icmp: [
{ value: 'type', label: 'type (icmp type)', kind: 'enum' },
{ value: 'code', label: 'code (icmp code)', kind: 'enum' },
],
/* 4) Layer 2: Ethernet & VLAN */
ether: [
{ value: 'saddr', label: 'saddr (src MAC)', kind: 'string' },
{ value: 'daddr', label: 'daddr (dst MAC)', kind: 'string' },
{ value: 'type', label: 'type (ethertype)', kind: 'enum' },
],
vlan: [
{ value: 'id', label: 'id (VLAN ID)', kind: 'number' },
// CFI/DEI and PCP exist but user's list specified only VLAN ID; add PCP & DEI as optional helpers:
{ value: 'pcp', label: 'pcp (priority code point)', kind: 'number' },
{ value: 'cfi', label: 'cfi / DEI (drop eligible)', kind: 'number' },
],
/* leftovers and statements */
payload: [{ value: 'payload', label: 'payload(protocol.field)', kind: 'string' }],
raw: [{ value: 'raw', label: 'raw text', kind: 'string' }],
counter: [{ value: 'counter', label: 'counter', kind: 'string' }],
limit: [{ value: 'limit', label: 'limit (rate)', kind: 'string' }],
log: [{ value: 'log', label: 'log', kind: 'string' }],
nat: [
{ value: 'dnat', label: 'dnat to', kind: 'string' },
{ value: 'snat', label: 'snat to', kind: 'string' },
{ value: 'masquerade', label: 'masquerade', kind: 'string' },
],
queue: [{ value: 'queue', label: 'queue num', kind: 'string' }],
};
/* ----------------------
ENUM_VALUES (dropdown contents)
Keep these aligned with the user's canonical lists.
---------------------- */
const ENUM_VALUES: Record<string, string[]> = {
l4proto: ['tcp', 'udp', 'icmp', 'icmpv6', 'igmp', 'esp', 'ah'],
days: ['Monday', 'Tuesday', 'Wednesday', 'Thursday', 'Friday', 'Saturday', 'Sunday'],
pkttype: ['unicast', 'multicast', 'broadcast', 'other'],
ct_state: ['new', 'established', 'related', 'invalid', 'untracked'],
ct_direction: ['original', 'reply'],
// ICMP message *types* (used for e.g. echo-request/echo-reply)
icmp_types: ['echo-request', 'echo-reply', 'destination-unreachable'],
// IPv4 reject *reasons* (ICMPv4 codes / textual reasons used with `reject with icmp type <reason>`)
icmpv4_reasons: [
'net-unreachable',
'host-unreachable',
'prot-unreachable',
'port-unreachable', // default
'net-prohibited',
'host-prohibited',
'admin-prohibited',
],
// IPv6 reject reasons (ICMPv6 textual reasons)
icmpv6_reasons: ['no-route', 'admin-prohibited', 'addr-unreachable', 'port-unreachable'],
dscp_values: [
'cs0',
'cs1',
'cs2',
'cs3',
'cs4',
'cs5',
'cs6',
'cs7',
'af11',
'af12',
'af13',
'af21',
'af22',
'af23',
'af31',
'af32',
'af33',
'af41',
'af42',
'af43',
'ef',
],
tcp_flags: ['fin', 'syn', 'rst', 'psh', 'ack', 'urg', 'ece', 'cwr'],
ethertypes: ['ip', 'ip6', 'arp', 'vlan', 'loopback'],
// top-level reject types used in select control. Note `icmpv6` spelled out.
reject_types: ['icmp', 'icmpv6', 'icmpx', 'tcp-reset'],
};
/* ----------------------
tokenToText: produce nft textual representation from token value
(keeps command generation consistent with the UI)
---------------------- */
function tokenToText(token: any): string {
if (!token || !token.type) return '';
const t = token.type as TokenType;
const d = token.data || {};
// META
if (t === 'meta') {
const f = d.field;
if (!f) return '';
// special formatting: meta l4proto <proto>
if (f === 'l4proto') {
return `meta l4proto ${String(d.value ?? '')}`.trim();
}
if (f === 'iifname' || f === 'oifname') {
return `meta ${f} ${String(d.value ?? '')}`.trim();
}
if (f === 'day') {
return `meta day ${String(d.value ?? '')}`.trim();
}
if (f === 'hour') {
return `meta hour ${String(d.value ?? '')}`.trim();
}
if (f === 'pkttype') {
return `meta pkttype ${String(d.value ?? '')}`.trim();
}
if (f === 'mark') {
return `meta mark ${String(d.value ?? '')}`.trim();
}
if (f === 'skuid' || f === 'skgid') {
return `meta ${f} ${String(d.value ?? '')}`.trim();
}
return `meta ${f} ${String(d.value ?? '')}`.trim();
}
// CT
if (t === 'ct') {
const f = d.field;
if (!f) return '';
return `ct ${f} ${String(d.value ?? '')}`.trim();
}
// IP/IPv6
if (t === 'ip' || t === 'ip6') {
const f = d.field;
if (!f) return '';
// saddr/daddr: allow CIDR/list/range raw text
return `${t} ${f} ${String(d.value ?? '')}`.trim();
}
// Transport protocols
if (t === 'tcp' || t === 'udp') {
const f = d.field;
if (!f) return t;
if (f === 'dport' || f === 'sport') {
return `${t} ${f} ${String(d.value ?? '')}`.trim();
}
if (f === 'flags') {
// flags could be array or comma-separated
const vals = Array.isArray(d.value)
? d.value
: String(d.value ?? '')
.split(',')
.map((s: string) => s.trim())
.filter(Boolean);
if (vals.length === 0) return t;
// render as: tcp flags { syn, ack }
return `${t} flags { ${vals.join(', ')} }`;
}
return `${t} ${f} ${String(d.value ?? '')}`.trim();
}
if (t === 'icmp') {
const f = d.field;
if (!f) return 'icmp';
return `icmp ${f} ${String(d.value ?? '')}`.trim();
}
// ETHER
if (t === 'ether') {
const f = d.field;
if (!f) return '';
return `ether ${f} ${String(d.value ?? '')}`.trim();
}
// VLAN
if (t === 'vlan') {
const f = d.field;
if (!f) return 'vlan';
return `vlan ${f} ${String(d.value ?? '')}`.trim();
}
// Statements
if (t === 'counter') {
if (d.packets || d.bytes) {
return `counter${d.packets ? ` packets ${d.packets}` : ''}${d.bytes ? ` bytes ${d.bytes}` : ''}`.trim();
}
return 'counter';
}
if (t === 'limit') {
const r = d.rate ?? d.value;
return r ? `limit rate ${r}` : 'limit';
}
if (t === 'log') {
const parts: string[] = [];
if (d.level) parts.push(`level ${d.level}`);
if (d.group) parts.push(`group ${d.group}`);
if (d.snaplen) parts.push(`snaplen ${d.snaplen}`);
if (d.prefix) parts.push(`prefix "${d.prefix}"`);
return parts.length ? `log ${parts.join(' ')}` : 'log';
}
if (t === 'nat') {
if (d.kind === 'dnat' && d.to) return `dnat to ${d.to}`;
if (d.kind === 'snat' && d.to) return `snat to ${d.to}`;
if (d.kind === 'masquerade') return d.to ? `masquerade to ${d.to}` : 'masquerade';
return 'nat';
}
if (t === 'queue') {
if (d.num) {
// allow optional extra token words following queue num, e.g. "queue num 1 bypass"
const extra = d.extra ? ` ${String(d.extra)}` : '';
return `queue num ${d.num}${extra}`.trim();
}
return 'queue';
}
if (t === 'raw') {
return String(d.text ?? '').trim();
}
if (t === 'payload') {
if (d.value) return `payload(${d.value})`;
return 'payload';
}
return '';
}
/* ----------------------
generateCommandFromValues (build textual + final nft add/insert)
---------------------- */
function generateCommandFromValues(values: FormValues) {
const tokens = Array.isArray(values.tokens) ? values.tokens : [];
const parts: string[] = [];
for (const t of tokens) {
const txt = tokenToText(t);
if (txt) parts.push(txt);
}
if (values.advanced && typeof values.advanced === 'string' && values.advanced.trim() !== '') {
parts.push(values.advanced.trim());
}
// Build queue text for NFQUEUE action or queue token
if (values.action === 'nfqueue' || values.action === 'queue') {
const qnum = values.nfqueue ?? values.queue ?? 1;
const bypass = values.nfqueue_bypass ? ' bypass' : '';
const queueText = `queue num ${Number(qnum)}${bypass}`;
const combined = parts.join(' ');
if (!/\bqueue(?:\s+num)?\b/i.test(combined)) {
parts.push(queueText);
} else {
for (let i = 0; i < parts.length; i++) {
if (/\bqueue(?:\s+num)?\b/i.test(parts[i])) {
parts[i] = queueText;
break;
}
}
}
}
// Build action/reject/nfqueue textual suffix
let actionText: string | null = null;
if (values.action === 'accept' || values.action === 'drop') {
actionText = values.action;
} else if (values.action === 'reject') {
// reject requires a rejectType (form enforces it)
const rtype = values.rejectType;
if (!rtype) {
actionText = 'reject'; // fallback, though form validation should prevent this
} else if (rtype === 'tcp-reset') {
// nft "reject with tcp reset"
actionText = 'reject with tcp reset';
} else if (rtype === 'icmp') {
// IPv4: "reject with icmp type <reason>"
const reason = values.rejectIcmpReason || '';
actionText = reason ? `reject with icmp type ${reason}` : 'reject';
} else if (rtype === 'icmpv6') {
// IPv6: "reject with icmpv6 type <reason>"
const reason = values.rejectIcmp6Reason || '';
actionText = reason ? `reject with icmpv6 type ${reason}` : 'reject';
} else if (rtype === 'icmpx') {
// inet family abstraction (icmpx)
const reason = values.rejectIcmpxReason || '';
actionText = reason ? `reject with icmpx type ${reason}` : 'reject';
} else {
actionText = 'reject';
}
} else if (values.action === 'nfqueue') {
// NFQUEUE action is represented by queue token above; no extra action verb
actionText = null;
}
const textual = (parts.join(' ') + (actionText ? ` ${actionText}` : '')).trim();
const tableSelect = values.tableSelect;
const chain = values.chainSelect || 'input';
const [family = 'inet', table = 'filter'] = tableSelect ? String(tableSelect).split(':') : ['inet', 'filter'];
const before = values.insertBeforeHandle;
const hasBefore = before != null && String(before) !== '';
const verb = hasBefore ? 'insert' : 'add';
const positionPart = hasBefore ? ` position ${before}` : '';
const cmd = `${verb} rule ${family} ${table} ${chain}${positionPart} ${textual}`.replace(/\s+/g, ' ').trim();
return { cmd, textual, position: hasBefore ? Number(before) : undefined };
}
/* -------------------------
Component
------------------------- */
interface RuleBuilderProps {
onCreated?: () => Promise<void> | void;
tables?: TableOut[] | null;
rulesLoading?: boolean;
rulesError?: string | null;
refreshRules?: () => Promise<void>;
onRulesChange?: (tables: TableOut[]) => void;
}
export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps) => {
const [form] = Form.useForm();
const [cmdPreview, setCmdPreview] = useState('');
const [refreshing, setRefreshing] = useState(false);
const [loading, setLoading] = useState(false);
const tableOptions = useMemo(
() => (props.tables || []).map((t) => ({ value: `${t.family}:${t.name}`, label: `${t.family}:${t.name}` })),
[props.tables],
);
const noTables = !(props.tables && props.tables.length > 0);
const [insertBeforeOptions, setInsertBeforeOptions] = useState<Array<{ value: any; label: string }>>([]);
const updateInsertOptions = useCallback(() => {
const ts = form.getFieldValue('tableSelect');
const cs = form.getFieldValue('chainSelect');
if (!ts || !cs) {
setInsertBeforeOptions([]);
return;
}
const [family, table] = String(ts).split(':');
const tbl = props.tables?.find((t) => t.family === family && t.name === table);
if (!tbl) {
setInsertBeforeOptions([]);
return;
}
const ch = (tbl.chains || []).find((c: ChainOut) => c.name === cs);
if (!ch || !Array.isArray(ch.rules)) {
setInsertBeforeOptions([]);
return;
}
const opts = ch.rules
.filter((r: RuleOut) => r && r.handle != null)
.map((r: RuleOut) => ({
value: r.handle,
label: `#${r.handle} — ${r.text ?? (typeof r.expr === 'string' ? r.expr : JSON.stringify(r.expr || {}).slice(0, 120))}`,
}));
setInsertBeforeOptions(opts);
}, [form, props.tables]);
const previewTimerRef = useRef<number | null>(null);
const schedulePreviewUpdate = useCallback(() => {
if (previewTimerRef.current) window.clearTimeout(previewTimerRef.current);
previewTimerRef.current = window.setTimeout(() => {
const v = form.getFieldsValue();
const { cmd } = generateCommandFromValues(v);
setCmdPreview(cmd);
previewTimerRef.current = null;
}, 40);
}, [form]);
useEffect(() => {
if (tableOptions.length > 0) {
const first = tableOptions[0].value;
form.setFieldsValue({
tableSelect: first,
action: 'drop',
nfqueue: 1,
nfqueue_bypass: false,
tokens: [],
});
const [f, n] = String(first).split(':');
const tbl = props.tables?.find((t) => t.family === f && t.name === n);
if (tbl && tbl.chains && tbl.chains.length > 0) {
form.setFieldsValue({ chainSelect: tbl.chains[0].name });
} else {
form.setFieldsValue({ chainSelect: undefined });
}
setTimeout(() => {
updateInsertOptions();
schedulePreviewUpdate();
}, 0);
} else {
form.setFieldsValue({
action: 'drop',
nfqueue: 1,
nfqueue_bypass: false,
tableSelect: undefined,
chainSelect: undefined,
tokens: [],
});
setInsertBeforeOptions([]);
setTimeout(() => schedulePreviewUpdate(), 0);
}
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [props.tables, tableOptions.length]);
useEffect(() => {
const cur = form.getFieldsValue();
if (cur.nfqueue == null) form.setFieldsValue({ nfqueue: 1 });
schedulePreviewUpdate();
// eslint-disable-next-line react-hooks/exhaustive-deps
}, []);
const onValuesChange = useCallback(
(_: any, allValues: FormValues) => {
if (allValues.action === 'nfqueue' && (allValues.nfqueue == null || allValues.nfqueue === '')) {
form.setFieldsValue({ nfqueue: 1 });
allValues.nfqueue = 1;
}
const ts = allValues.tableSelect;
if (ts) {
const [f, n] = String(ts).split(':');
const tbl = props.tables?.find((t) => t.family === f && t.name === n);
if (tbl) {
if (tbl.chains && tbl.chains.length > 0) {
if (!allValues.chainSelect) form.setFieldsValue({ chainSelect: tbl.chains[0].name });
} else {
form.setFieldsValue({ chainSelect: undefined });
}
}
}
updateInsertOptions();
schedulePreviewUpdate();
},
[form, props.tables, updateInsertOptions, schedulePreviewUpdate],
);
const handleCreate = useCallback(
async (values: FormValues) => {
try {
const validated = await form.validateFields();
const { cmd } = generateCommandFromValues(validated);
Modal.confirm({
title: 'Run raw nft command',
content: (
<div>
<Text>
About to run nft command in <b>{String(validated.tableSelect ?? 'inet:filter')}</b> (see preview).
</Text>
<Divider />
<Text strong>Command:</Text>
<pre style={{ whiteSpace: 'pre-wrap', marginTop: 8 }}>{cmd}</pre>
</div>
),
okText: 'Run',
onOk: async () => {
setLoading(true);
try {
const out: ExecResult = await execFirewallRaw(cmd);
const stderrText = out?.stderr ? String(out.stderr).trim() : '';
if (stderrText) {
notification.error({
message: 'Command produced Error',
description: stderrText,
});
} else if (out && (out.rc === 0 || out.rc === -1)) {
notification.success({
message: 'Command executed successfully',
});
if (props.refreshRules) await props.refreshRules();
if (props.onCreated) await props.onCreated();
} else {
const info = out
? `rc:${out.rc}` +
(out.stdout ? ` stdout:${out.stdout}` : '') +
(out.stderr ? ` stderr:${out.stderr}` : '')
: 'unknown result';
notification.error({
message: 'Command failed',
description: info,
});
}
} catch (err: any) {
notification.error({
message: 'Execution failed',
description: err?.message ?? String(err),
});
} finally {
setLoading(false);
}
},
});
} catch (err) {
schedulePreviewUpdate();
}
},
[form, props.refreshRules, props.onCreated, schedulePreviewUpdate],
);
const chainOptions = useMemo(() => {
const ts = form.getFieldValue('tableSelect');
if (!ts) return [];
const [f, n] = String(ts).split(':');
const tbl = props.tables?.find((t) => t.family === f && t.name === n);
if (!tbl) return [];
return tbl.chains.map((c) => (
<Option key={c.name} value={c.name}>
{c.name}
</Option>
));
}, [form, props.tables]);
const handleRefresh = useCallback(async () => {
setRefreshing(true);
try {
if (props.refreshRules) {
await props.refreshRules();
message.success('Rules refresh requested');
} else {
message.info('No refresh function provided by parent.');
}
} catch (err) {
console.warn('refresh failed', err);
message.error('Refresh failed');
} finally {
updateInsertOptions();
setRefreshing(false);
}
}, [props.refreshRules, updateInsertOptions]);
/* helper styles */
const tokenRowStyle: React.CSSProperties = {
display: 'flex',
gap: 8,
alignItems: 'center',
flexWrap: 'nowrap',
width: '100%',
};
const leftControlsStyle: React.CSSProperties = {
display: 'flex',
gap: 8,
alignItems: 'center',
minWidth: 72,
flex: '0 0 72px',
};
const typeSelectStyle: React.CSSProperties = { minWidth: 180, maxWidth: 260, flex: '0 0 220px' };
const fieldSelectStyle: React.CSSProperties = { minWidth: 160, maxWidth: 260, flex: '0 0 220px' };
const valueInputStyle: React.CSSProperties = { minWidth: 120, flex: '1 1 240px', maxWidth: '60%' };
const actionControlsStyle: React.CSSProperties = {
minWidth: 96,
flex: '0 0 96px',
display: 'flex',
justifyContent: 'flex-end',
};
return (
<Card title="Firewall Rule Builder" extra>
<Form
layout="vertical"
form={form}
initialValues={{
action: 'drop',
nfqueue: 1,
nfqueue_bypass: false,
tableSelect: tableOptions.length > 0 ? tableOptions[0].value : undefined,
tokens: [],
}}
onFinish={handleCreate}
onValuesChange={onValuesChange}
>
{/* Table / chain */}
<Row gutter={16} align="middle"> <Row gutter={16} align="middle">
<Col xs={24} sm={12}> <Col xs={24} sm={12}>
<Form.Item name="tableSelect" label="Table (family:name)" rules={[{ required: true }]}> <Form.Item name="tableSelect" label="Table (family:name)" rules={[{ required: true }]}>
@@ -717,7 +71,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</Col> </Col>
</Row> </Row>
{/* Insert before */}
<Row gutter={16} align="middle"> <Row gutter={16} align="middle">
<Col xs={24} sm={12}> <Col xs={24} sm={12}>
<Form.Item <Form.Item
@@ -750,7 +104,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
<Divider /> <Divider />
{/* Token builder header + add control */}
<Row align="middle" justify="space-between" style={{ marginBottom: 8 }}> <Row align="middle" justify="space-between" style={{ marginBottom: 8 }}>
<Col> <Col>
<Text strong>Token builder</Text> <Text strong>Token builder</Text>
@@ -766,7 +120,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
<Select <Select
placeholder="Add token..." placeholder="Add token..."
onSelect={(val: TokenType) => { onSelect={(val: TokenType) => {
// sensible defaults per token type
const defaultData = const defaultData =
val === 'meta' val === 'meta'
? { field: 'iifname', value: '' } ? { field: 'iifname', value: '' }
@@ -823,7 +176,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
<Divider /> <Divider />
{/* Tokens Form.List rendering */}
<Form.List name="tokens"> <Form.List name="tokens">
{(fields, { remove, move }) => {(fields, { remove, move }) =>
fields.length === 0 ? ( fields.length === 0 ? (
@@ -842,7 +195,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</div> </div>
<div style={{ display: 'flex', gap: 8, alignItems: 'center', flex: 1, minWidth: 0 }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', flex: 1, minWidth: 0 }}>
{/* Token type select */}
<Form.Item name={[field.name, 'type']} style={{ marginBottom: 0 }}> <Form.Item name={[field.name, 'type']} style={{ marginBottom: 0 }}>
<Select style={typeSelectStyle}> <Select style={typeSelectStyle}>
{Object.keys(TOKEN_FIELD_OPTIONS).map((k) => ( {Object.keys(TOKEN_FIELD_OPTIONS).map((k) => (
@@ -853,7 +206,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</Select> </Select>
</Form.Item> </Form.Item>
{/* Token field + value UI (depends on token type and subfield) */}
<Form.Item <Form.Item
shouldUpdate={(prev, cur) => shouldUpdate={(prev, cur) =>
prev.tokens?.[field.name]?.type !== cur.tokens?.[field.name]?.type || prev.tokens?.[field.name]?.type !== cur.tokens?.[field.name]?.type ||
@@ -865,7 +218,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
const tokenType = form.getFieldValue(['tokens', field.name, 'type']) as TokenType | undefined; const tokenType = form.getFieldValue(['tokens', field.name, 'type']) as TokenType | undefined;
const options = tokenType ? TOKEN_FIELD_OPTIONS[tokenType] || [] : []; const options = tokenType ? TOKEN_FIELD_OPTIONS[tokenType] || [] : [];
// COUNTER special-case
if (tokenType === 'counter') { if (tokenType === 'counter') {
return ( return (
<div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}>
@@ -885,7 +237,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// LIMIT special-case
if (tokenType === 'limit') { if (tokenType === 'limit') {
return ( return (
<div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}>
@@ -903,7 +254,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// LOG special-case
if (tokenType === 'log') { if (tokenType === 'log') {
return ( return (
<div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}>
@@ -941,7 +291,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// QUEUE special-case inside token list (separate from NFQUEUE action)
if (tokenType === 'queue') { if (tokenType === 'queue') {
return ( return (
<div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}>
@@ -961,7 +310,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// NAT special-case
if (tokenType === 'nat') { if (tokenType === 'nat') {
return ( return (
<div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}> <div style={{ display: 'flex', gap: 8, alignItems: 'center', width: '100%' }}>
@@ -983,7 +331,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// Generic tokens with subfield dropdown
if (options.length > 0) { if (options.length > 0) {
return ( return (
<div <div
@@ -1017,9 +364,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
const meta = opts.find((o) => o.value === selField); const meta = opts.find((o) => o.value === selField);
const kind = meta?.kind ?? 'string'; const kind = meta?.kind ?? 'string';
/* --- Field-specific UIs & helpers (placeholders + explanatory text) --- */
// STRING typed helpers for interface names
if (tType === 'meta' && (selField === 'iifname' || selField === 'oifname')) { if (tType === 'meta' && (selField === 'iifname' || selField === 'oifname')) {
return ( return (
<div> <div>
@@ -1034,7 +379,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// L4PROTO dropdown for meta.l4proto
if (tType === 'meta' && selField === 'l4proto') { if (tType === 'meta' && selField === 'l4proto') {
return ( return (
<div> <div>
@@ -1052,7 +396,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// Day of week (meta.day)
if (tType === 'meta' && selField === 'day') { if (tType === 'meta' && selField === 'day') {
return ( return (
<div> <div>
@@ -1072,7 +415,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// Hour range (meta.hour) — free text but show placeholder/range hint
if (tType === 'meta' && selField === 'hour') { if (tType === 'meta' && selField === 'hour') {
return ( return (
<div> <div>
@@ -1087,7 +429,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// Packet type (meta.pkttype)
if (tType === 'meta' && selField === 'pkttype') { if (tType === 'meta' && selField === 'pkttype') {
return ( return (
<div> <div>
@@ -1107,7 +448,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// Packet/conn mark
if ((tType === 'meta' || tType === 'ct') && selField === 'mark') { if ((tType === 'meta' || tType === 'ct') && selField === 'mark') {
return ( return (
<div> <div>
@@ -1121,7 +461,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// skuid / skgid
if (tType === 'meta' && (selField === 'skuid' || selField === 'skgid')) { if (tType === 'meta' && (selField === 'skuid' || selField === 'skgid')) {
return ( return (
<div> <div>
@@ -1140,7 +479,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// CT state
if (tType === 'ct' && selField === 'state') { if (tType === 'ct' && selField === 'state') {
return ( return (
<div> <div>
@@ -1165,7 +503,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// CT direction
if (tType === 'ct' && selField === 'direction') { if (tType === 'ct' && selField === 'direction') {
return ( return (
<div> <div>
@@ -1185,8 +522,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
/* --- IP / IP6 address helpers --- */ if (
if (
(tType === 'ip' || tType === 'ip6') && (tType === 'ip' || tType === 'ip6') &&
(selField === 'saddr' || selField === 'daddr') (selField === 'saddr' || selField === 'daddr')
) { ) {
@@ -1214,7 +550,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// protocol / nexthdr / ip.protocol (L4 protocol): show l4proto list
if ( if (
(tType === 'ip' && selField === 'protocol') || (tType === 'ip' && selField === 'protocol') ||
(tType === 'ip6' && selField === 'nexthdr') (tType === 'ip6' && selField === 'nexthdr')
@@ -1238,7 +573,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// DSCP
if ((tType === 'ip' || tType === 'ip6') && selField === 'dscp') { if ((tType === 'ip' || tType === 'ip6') && selField === 'dscp') {
return ( return (
<div> <div>
@@ -1258,7 +592,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// TTL / hoplimit numeric
if ( if (
(tType === 'ip' && selField === 'ttl') || (tType === 'ip' && selField === 'ttl') ||
(tType === 'ip6' && selField === 'hoplimit') (tType === 'ip6' && selField === 'hoplimit')
@@ -1278,7 +611,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// IP fragment bits (frag-off) — single string placeholder
if (tType === 'ip' && selField === 'frag-off') { if (tType === 'ip' && selField === 'frag-off') {
return ( return (
<div> <div>
@@ -1292,9 +624,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
/* --- Transport: TCP/UDP/ICMP --- */
// Ports: allow numeric or service name
if ( if (
(tType === 'tcp' || tType === 'udp') && (tType === 'tcp' || tType === 'udp') &&
(selField === 'dport' || selField === 'sport') (selField === 'dport' || selField === 'sport')
@@ -1311,7 +641,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// TCP flags multi-select
if (tType === 'tcp' && selField === 'flags') { if (tType === 'tcp' && selField === 'flags') {
return ( return (
<div> <div>
@@ -1335,7 +664,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// ICMP type/code dropdowns
if (tType === 'icmp' && selField === 'type') { if (tType === 'icmp' && selField === 'type') {
return ( return (
<div> <div>
@@ -1374,8 +702,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
/* --- Layer 2: Ethernet / VLAN --- */
if (tType === 'ether') { if (tType === 'ether') {
if (selField === 'saddr' || selField === 'daddr') { if (selField === 'saddr' || selField === 'daddr') {
return ( return (
@@ -1408,7 +735,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
} }
} }
// VLAN ID / PCP / CFI
if (tType === 'vlan') { if (tType === 'vlan') {
if (selField === 'id') { if (selField === 'id') {
return ( return (
@@ -1459,8 +785,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
} }
} }
/* --- Payload / default free text input --- */ if (tType === 'payload' || (kind === 'string' && !selField)) {
if (tType === 'payload' || (kind === 'string' && !selField)) {
return ( return (
<div> <div>
<Form.Item name={[field.name, 'data', 'value']} style={{ margin: 0 }}> <Form.Item name={[field.name, 'data', 'value']} style={{ margin: 0 }}>
@@ -1474,7 +799,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
); );
} }
// Default fallback free-text with helpful examples
return ( return (
<div> <div>
<Form.Item name={[field.name, 'data', 'value']} style={{ margin: 0 }}> <Form.Item name={[field.name, 'data', 'value']} style={{ margin: 0 }}>
@@ -1515,8 +839,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
<Divider /> <Divider />
{/* Action + NFQUEUE + Reject options: render action radios, then render
reject subform and nfqueue subform directly under it (same column) */}
<Row gutter={16} align="top"> <Row gutter={16} align="top">
<Col xs={24} sm={12}> <Col xs={24} sm={12}>
<Form.Item name="action" label="Action / verdict" rules={[{ required: true }]}> <Form.Item name="action" label="Action / verdict" rules={[{ required: true }]}>
@@ -1528,7 +851,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</Radio.Group> </Radio.Group>
</Form.Item> </Form.Item>
{/* Reject options (render under radios, same column) */}
<Form.Item shouldUpdate={(prev, cur) => prev.action !== cur.action} noStyle> <Form.Item shouldUpdate={(prev, cur) => prev.action !== cur.action} noStyle>
{() => {() =>
form.getFieldValue('action') === 'reject' ? ( form.getFieldValue('action') === 'reject' ? (
@@ -1547,7 +870,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</Select> </Select>
</Form.Item> </Form.Item>
{/* IPv4 reject reasons */}
<Form.Item shouldUpdate={(prev, cur) => prev.rejectType !== cur.rejectType} noStyle> <Form.Item shouldUpdate={(prev, cur) => prev.rejectType !== cur.rejectType} noStyle>
{() => {() =>
form.getFieldValue('rejectType') === 'icmp' ? ( form.getFieldValue('rejectType') === 'icmp' ? (
@@ -1572,7 +895,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
} }
</Form.Item> </Form.Item>
{/* IPv6 reject reasons */}
<Form.Item shouldUpdate={(prev, cur) => prev.rejectType !== cur.rejectType} noStyle> <Form.Item shouldUpdate={(prev, cur) => prev.rejectType !== cur.rejectType} noStyle>
{() => {() =>
form.getFieldValue('rejectType') === 'icmpv6' ? ( form.getFieldValue('rejectType') === 'icmpv6' ? (
@@ -1597,7 +920,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
} }
</Form.Item> </Form.Item>
{/* icmpx (inet) */}
<Form.Item shouldUpdate={(prev, cur) => prev.rejectType !== cur.rejectType} noStyle> <Form.Item shouldUpdate={(prev, cur) => prev.rejectType !== cur.rejectType} noStyle>
{() => {() =>
form.getFieldValue('rejectType') === 'icmpx' ? ( form.getFieldValue('rejectType') === 'icmpx' ? (
@@ -1628,7 +951,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
} }
</Form.Item> </Form.Item>
{/* NFQUEUE options (now rendered under radios in same column) */}
<Form.Item shouldUpdate={(prev, cur) => prev.action !== cur.action} noStyle> <Form.Item shouldUpdate={(prev, cur) => prev.action !== cur.action} noStyle>
{() => {() =>
form.getFieldValue('action') === 'nfqueue' ? ( form.getFieldValue('action') === 'nfqueue' ? (
@@ -1648,7 +971,6 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
<Form.Item name="nfqueue_bypass" valuePropName="checked"> <Form.Item name="nfqueue_bypass" valuePropName="checked">
<Checkbox <Checkbox
onChange={() => { onChange={() => {
// update preview immediately
schedulePreviewUpdate(); schedulePreviewUpdate();
}} }}
> >
@@ -1665,7 +987,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</Form.Item> </Form.Item>
</Col> </Col>
{/* right column is free for notes / quick helpers */}
<Col xs={24} sm={12}> <Col xs={24} sm={12}>
<Text type="secondary"> <Text type="secondary">
Use NFQUEUE to hand packets to userspace. Full reject support requires kernel &gt;= 3.18 — when using Use NFQUEUE to hand packets to userspace. Full reject support requires kernel &gt;= 3.18 — when using
@@ -1674,7 +996,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
</Col> </Col>
</Row> </Row>
{/* advanced text */}
<Row> <Row>
<Col span={24}> <Col span={24}>
<Form.Item name="advanced" label="Advanced (optional)"> <Form.Item name="advanced" label="Advanced (optional)">
@@ -1688,7 +1010,7 @@ export const RuleBuilder: React.FC<RuleBuilderProps> = (props: RuleBuilderProps)
<Divider /> <Divider />
{/* preview + run */}
<Form.Item> <Form.Item>
<div style={{ display: 'flex', gap: 12, alignItems: 'center', width: '100%' }}> <div style={{ display: 'flex', gap: 12, alignItems: 'center', width: '100%' }}>
<div style={{ flex: 1 }}> <div style={{ flex: 1 }}>

View File

@@ -1,4 +1,3 @@
// src/components/FirewallTables.tsx
import { ArrowDownOutlined, DeleteOutlined, ReloadOutlined } from '@ant-design/icons'; import { ArrowDownOutlined, DeleteOutlined, ReloadOutlined } from '@ant-design/icons';
import { Alert, Button, Card, Divider, Modal, notification, Space, Spin, Table, Typography } from 'antd'; import { Alert, Button, Card, Divider, Modal, notification, Space, Spin, Table, Typography } from 'antd';
import { ColumnsType } from 'antd/lib/table'; import { ColumnsType } from 'antd/lib/table';
@@ -10,7 +9,6 @@ import FirewallAddTableModal from './FireWallAddTableModal';
const { Paragraph, Text, Title } = Typography; const { Paragraph, Text, Title } = Typography;
/* ---------- Helpers ---------- */
function renderRuleFriendly(rule: RuleOut | any): string { function renderRuleFriendly(rule: RuleOut | any): string {
if (rule?.text && typeof rule.text === 'string' && rule.text.trim() !== '') return rule.text; if (rule?.text && typeof rule.text === 'string' && rule.text.trim() !== '') return rule.text;
@@ -103,7 +101,6 @@ function renderRuleFriendly(rule: RuleOut | any): string {
try { try {
return JSON.stringify(rule.expr, (_k, v) => (v === undefined ? null : v)).slice(0, 500); return JSON.stringify(rule.expr, (_k, v) => (v === undefined ? null : v)).slice(0, 500);
} catch { } catch {
// fallthrough
} }
} }
@@ -120,22 +117,18 @@ function isSuccessRc(out?: ExecResult | null): boolean {
return out.rc === 0 || out.rc === -1; return out.rc === 0 || out.rc === -1;
} }
/* ---------- Props ---------- */
type Props = { type Props = {
tables: TableOut[]; // passed from parent tables: TableOut[]; // passed from parent
error?: Error | null; error?: Error | null;
refreshRules: () => Promise<void>; // trigger to re-fetch ruleset refreshRules: () => Promise<void>; // trigger to re-fetch ruleset
}; };
/* ---------- Component ---------- */
export default function FirewallTables({ tables, error, refreshRules: refresh }: Props): ReactElement { export default function FirewallTables({ tables, error, refreshRules: refresh }: Props): ReactElement {
// local UI state, non-persistent
const [refreshing, setRefreshing] = useState(false); const [refreshing, setRefreshing] = useState(false);
const [isOpenTableCreatorModal, setIsOpenTableCreatorModal] = useState(false); const [isOpenTableCreatorModal, setIsOpenTableCreatorModal] = useState(false);
const [isOpenChainCreatorModal, setIsOpenChainCreatorModal] = useState(false); const [isOpenChainCreatorModal, setIsOpenChainCreatorModal] = useState(false);
// run raw nft commands sequentially and collect results (used for delete ops etc.)
const runCommands = useCallback(async (cmds: string[]) => { const runCommands = useCallback(async (cmds: string[]) => {
const acc: CmdResult[] = []; const acc: CmdResult[] = [];
for (const cmd of cmds) { for (const cmd of cmds) {
@@ -150,7 +143,6 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
return acc; return acc;
}, []); }, []);
// Delete rule
const handleDeleteRule = useCallback( const handleDeleteRule = useCallback(
async (family: string | null | undefined, table: string, chain: string, handle: number | string) => { async (family: string | null | undefined, table: string, chain: string, handle: number | string) => {
const cmd = `delete rule ${family ?? 'inet'} ${table} ${chain} handle ${handle}`; const cmd = `delete rule ${family ?? 'inet'} ${table} ${chain} handle ${handle}`;
@@ -175,12 +167,9 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
} catch (err: any) { } catch (err: any) {
notification.error({ message: 'Delete failed', description: err?.message ?? String(err) }); notification.error({ message: 'Delete failed', description: err?.message ?? String(err) });
} finally { } finally {
// auto-refresh after change (no refresh notification shown here)
try { try {
await refresh(); await refresh();
} catch { } catch
/* ignore */
}
} }
}, },
}); });
@@ -188,7 +177,6 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
[runCommands, refresh], [runCommands, refresh],
); );
// Delete chain
const handleDeleteChain = useCallback( const handleDeleteChain = useCallback(
async (family: string | null | undefined, table: string, chain: string) => { async (family: string | null | undefined, table: string, chain: string) => {
const cmd = `delete chain ${family ?? 'inet'} ${table} ${chain}`; const cmd = `delete chain ${family ?? 'inet'} ${table} ${chain}`;
@@ -215,12 +203,9 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
} catch (err: any) { } catch (err: any) {
notification.error({ message: 'Chain deletion failed', description: err?.message ?? String(err) }); notification.error({ message: 'Chain deletion failed', description: err?.message ?? String(err) });
} finally { } finally {
// auto-refresh after change (no notification)
try { try {
await refresh(); await refresh();
} catch { } catch
/* ignore */
}
} }
}, },
}); });
@@ -228,7 +213,6 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
[runCommands, refresh], [runCommands, refresh],
); );
// Delete table
const handleDeleteTable = useCallback( const handleDeleteTable = useCallback(
async (family: string | null | undefined, table: string) => { async (family: string | null | undefined, table: string) => {
const cmd = `delete table ${family ?? 'inet'} ${table}`; const cmd = `delete table ${family ?? 'inet'} ${table}`;
@@ -255,12 +239,9 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
} catch (err: any) { } catch (err: any) {
notification.error({ message: 'Table deletion failed', description: err?.message ?? String(err) }); notification.error({ message: 'Table deletion failed', description: err?.message ?? String(err) });
} finally { } finally {
// auto-refresh after change (no notification)
try { try {
await refresh(); await refresh();
} catch { } catch
/* ignore */
}
} }
}, },
}); });
@@ -268,12 +249,10 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
[runCommands, refresh], [runCommands, refresh],
); );
// manual refresh trigger (exposed to UI)
const handleRefresh = useCallback(async () => { const handleRefresh = useCallback(async () => {
setRefreshing(true); setRefreshing(true);
try { try {
await refresh(); await refresh();
// Only show notification when user pressed the refresh button
notification.success({ message: 'Ruleset refreshed' }); notification.success({ message: 'Ruleset refreshed' });
} catch (err: any) { } catch (err: any) {
notification.error({ message: 'Refresh failed', description: err?.message ?? String(err) }); notification.error({ message: 'Refresh failed', description: err?.message ?? String(err) });
@@ -289,14 +268,12 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
return ( return (
<> <>
{/* Pass onClose that accepts optional 'created' boolean. If the modal
calls onClose(true) we will auto-refresh (no refresh notification). */}
<FirewallAddTableModal <FirewallAddTableModal
open={isOpenTableCreatorModal} open={isOpenTableCreatorModal}
onClose={(created?: boolean) => { onClose={(created?: boolean) => {
setIsOpenTableCreatorModal(false); setIsOpenTableCreatorModal(false);
if (created) { if (created) {
// auto-refresh after create (no notification)
void refresh().catch(() => {}); void refresh().catch(() => {});
} }
}} }}
@@ -361,7 +338,7 @@ export default function FirewallTables({ tables, error, refreshRules: refresh }:
</div> </div>
} }
> >
{/* Chain modal: same optional 'created' signal */}
<FirewallAddChainModal <FirewallAddChainModal
open={isOpenChainCreatorModal} open={isOpenChainCreatorModal}
onClose={(created?: boolean) => { onClose={(created?: boolean) => {

View File

@@ -1,4 +1,3 @@
// src/pages/ScriptsManager.tsx
import { import {
DeleteOutlined, DeleteOutlined,
DownloadOutlined, DownloadOutlined,
@@ -65,49 +64,36 @@ const { Paragraph } = Typography;
*/ */
export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (code: string) => void }) { export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (code: string) => void }) {
// --------------------
// State
// --------------------
const [scripts, setScripts] = useState<ScriptWithStatus[]>([]); const [scripts, setScripts] = useState<ScriptWithStatus[]>([]);
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
// modals & editor values
const [uploadModalVisible, setUploadModalVisible] = useState(false); const [uploadModalVisible, setUploadModalVisible] = useState(false);
const [editorModalVisible, setEditorModalVisible] = useState(false); const [editorModalVisible, setEditorModalVisible] = useState(false);
const [editorValue, setEditorValue] = useState<string>(''); const [editorValue, setEditorValue] = useState<string>('');
const [currentEditingName, setCurrentEditingName] = useState<string | null>(null); const [currentEditingName, setCurrentEditingName] = useState<string | null>(null);
// inline requirements editor toggle + value for upload modal
const [useInlineReqEditor, setUseInlineReqEditor] = useState(false); const [useInlineReqEditor, setUseInlineReqEditor] = useState(false);
const [inlineReqValue, setInlineReqValue] = useState<string>(''); const [inlineReqValue, setInlineReqValue] = useState<string>('');
// pip output modal
const [pipModalVisible, setPipModalVisible] = useState(false); const [pipModalVisible, setPipModalVisible] = useState(false);
const [pipOutput, setPipOutput] = useState<{ stdout?: string; stderr?: string } | null>(null); const [pipOutput, setPipOutput] = useState<{ stdout?: string; stderr?: string } | null>(null);
// requirements editor modal (existing per-script editor)
const [reqModalVisible, setReqModalVisible] = useState(false); const [reqModalVisible, setReqModalVisible] = useState(false);
const [reqEditorValue, setReqEditorValue] = useState<string>(''); const [reqEditorValue, setReqEditorValue] = useState<string>('');
const [reqEditingName, setReqEditingName] = useState<string | null>(null); const [reqEditingName, setReqEditingName] = useState<string | null>(null);
// add-requirements modal (for scripts that have no requirements)
const [addReqModalVisible, setAddReqModalVisible] = useState(false); const [addReqModalVisible, setAddReqModalVisible] = useState(false);
const [addReqTarget, setAddReqTarget] = useState<string | null>(null); const [addReqTarget, setAddReqTarget] = useState<string | null>(null);
const [addReqUseInline, setAddReqUseInline] = useState(false); const [addReqUseInline, setAddReqUseInline] = useState(false);
const [addReqInlineValue, setAddReqInlineValue] = useState(''); const [addReqInlineValue, setAddReqInlineValue] = useState('');
const [addReqFile, setAddReqFile] = useState<File | null>(null); const [addReqFile, setAddReqFile] = useState<File | null>(null);
// enable modal
const [enableModalVisible, setEnableModalVisible] = useState(false); const [enableModalVisible, setEnableModalVisible] = useState(false);
const [enableTarget, setEnableTarget] = useState<string | null>(null); const [enableTarget, setEnableTarget] = useState<string | null>(null);
// forms
const [form] = Form.useForm(); const [form] = Form.useForm();
const [enableForm] = Form.useForm(); const [enableForm] = Form.useForm();
// --------------------
// Data fetch
// --------------------
const refreshAll = useCallback(async () => { const refreshAll = useCallback(async () => {
setLoading(true); setLoading(true);
try { try {
@@ -148,9 +134,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
refreshAll(); refreshAll();
}, [refreshAll]); }, [refreshAll]);
// --------------------
// Helpers
// --------------------
const getMappingsForScript = useCallback( const getMappingsForScript = useCallback(
(scriptName: string) => { (scriptName: string) => {
const s = scripts.find((x) => x.name === scriptName); const s = scripts.find((x) => x.name === scriptName);
@@ -171,7 +154,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
[scripts], [scripts],
); );
// Focus helper: tries to find a textarea inside the modal and put caret at the end.
const focusEditorInModal = useCallback((modalSelector = '.ant-modal') => { const focusEditorInModal = useCallback((modalSelector = '.ant-modal') => {
setTimeout(() => { setTimeout(() => {
try { try {
@@ -184,7 +166,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
ta.setSelectionRange(val.length, val.length); ta.setSelectionRange(val.length, val.length);
} }
} catch { } catch {
// ignore
} }
}, 80); }, 80);
}, []); }, []);
@@ -205,7 +186,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
if (addReqModalVisible) focusEditorInModal('.ant-modal'); if (addReqModalVisible) focusEditorInModal('.ant-modal');
}, [addReqModalVisible, focusEditorInModal]); }, [addReqModalVisible, focusEditorInModal]);
// show pip modal if pip output is present
const showPipIfPresent = useCallback( const showPipIfPresent = useCallback(
(resp?: ScriptUploadResponse | { pip?: { stdout?: string; stderr?: string } } | null) => { (resp?: ScriptUploadResponse | { pip?: { stdout?: string; stderr?: string } } | null) => {
const pip = (resp as any)?.pip; const pip = (resp as any)?.pip;
@@ -217,9 +197,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
[], [],
); );
// --------------------
// API Handlers
// --------------------
const handleDownload = useCallback(async (name: string) => { const handleDownload = useCallback(async (name: string) => {
try { try {
const blob = await downloadScript(name); const blob = await downloadScript(name);
@@ -358,9 +335,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
[refreshAll], [refreshAll],
); );
// --------------------
// Upload handler (prefers file; falls back to inline editors)
// --------------------
const handleUpload = useCallback( const handleUpload = useCallback(
async (formValues: any) => { async (formValues: any) => {
const { name, scriptFile, requirementsFile } = formValues; const { name, scriptFile, requirementsFile } = formValues;
@@ -369,10 +343,8 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
return; return;
} }
// script: prefer uploaded file; else use editorValue
const script = scriptFile || new Blob([editorValue], { type: 'text/x-python' }); const script = scriptFile || new Blob([editorValue], { type: 'text/x-python' });
// requirements: prefer uploaded file; else use inlineReqValue if toggle enabled; else null
let req: File | Blob | null = null; let req: File | Blob | null = null;
if (requirementsFile) { if (requirementsFile) {
req = requirementsFile; req = requirementsFile;
@@ -476,10 +448,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
[refreshAll], [refreshAll],
); );
// --------------------
// Upload form helpers
// --------------------
// Normalize Upload event to return single File object for Form storage
const normFile = (e: any) => { const normFile = (e: any) => {
if (!e) return undefined; if (!e) return undefined;
const list: UploadFile[] = e.fileList || []; const list: UploadFile[] = e.fileList || [];
@@ -488,7 +456,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
return (last as any).originFileObj ?? last; return (last as any).originFileObj ?? last;
}; };
// Icon-only button wrapped in Tooltip
function IconButtonTooltip({ function IconButtonTooltip({
title, title,
onClick, onClick,
@@ -509,9 +476,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
); );
} }
// --------------------
// Add Requirements (for scripts missing requirements)
// --------------------
const openAddRequirements = useCallback((scriptName: string) => { const openAddRequirements = useCallback((scriptName: string) => {
setAddReqTarget(scriptName); setAddReqTarget(scriptName);
setAddReqUseInline(false); setAddReqUseInline(false);
@@ -549,7 +513,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
} }
}, [addReqFile, addReqInlineValue, addReqTarget, refreshAll, showPipIfPresent]); }, [addReqFile, addReqInlineValue, addReqTarget, refreshAll, showPipIfPresent]);
// helper for Upload change inside add-req modal
const onAddReqUploadChange = useCallback((info: any) => { const onAddReqUploadChange = useCallback((info: any) => {
const list: UploadFile[] = info.fileList || []; const list: UploadFile[] = info.fileList || [];
if (list.length === 0) { if (list.length === 0) {
@@ -560,9 +523,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
setAddReqFile(last.originFileObj ?? last); setAddReqFile(last.originFileObj ?? last);
}, []); }, []);
// --------------------
// Table columns
// --------------------
const columns: ColumnsType<ScriptWithStatus> = useMemo( const columns: ColumnsType<ScriptWithStatus> = useMemo(
() => [ () => [
{ {
@@ -715,9 +675,6 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
[getMappingsForScript, nestedColumns], [getMappingsForScript, nestedColumns],
); );
// --------------------
// Render
// --------------------
return ( return (
<div> <div>
<Row justify="space-between" style={{ marginBottom: 12 }}> <Row justify="space-between" style={{ marginBottom: 12 }}>
@@ -740,7 +697,7 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
expandable={expandable} expandable={expandable}
/> />
{/* Upload Modal (wider) */}
<Modal <Modal
open={uploadModalVisible} open={uploadModalVisible}
title="Upload or create script" title="Upload or create script"
@@ -816,7 +773,7 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
</Form> </Form>
</Modal> </Modal>
{/* Editor Modal */}
<Modal <Modal
open={editorModalVisible} open={editorModalVisible}
title={currentEditingName ? `Editing — ${currentEditingName}` : 'Editor'} title={currentEditingName ? `Editing — ${currentEditingName}` : 'Editor'}
@@ -874,7 +831,7 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
<PythonEditor value={editorValue} onChange={setEditorValue} height={520} /> <PythonEditor value={editorValue} onChange={setEditorValue} height={520} />
</Modal> </Modal>
{/* Requirements Editor Modal (per-script) */}
<Modal <Modal
open={reqModalVisible} open={reqModalVisible}
title={reqEditingName ? `requirements.txt — ${reqEditingName}` : 'requirements.txt'} title={reqEditingName ? `requirements.txt — ${reqEditingName}` : 'requirements.txt'}
@@ -913,7 +870,7 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
<PythonEditor value={reqEditorValue} onChange={setReqEditorValue} height={420} /> <PythonEditor value={reqEditorValue} onChange={setReqEditorValue} height={420} />
</Modal> </Modal>
{/* Add Requirements Modal (for scripts without requirements) */}
<Modal <Modal
open={addReqModalVisible} open={addReqModalVisible}
title={addReqTarget ? `Add requirements — ${addReqTarget}` : 'Add requirements'} title={addReqTarget ? `Add requirements — ${addReqTarget}` : 'Add requirements'}
@@ -967,7 +924,7 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
</div> </div>
</Modal> </Modal>
{/* Pip output modal */}
<Modal <Modal
open={pipModalVisible} open={pipModalVisible}
title="pip install output" title="pip install output"
@@ -986,7 +943,7 @@ export default function ScriptsManager({ onOpenInEditor }: { onOpenInEditor?: (c
</div> </div>
</Modal> </Modal>
{/* Enable Modal */}
<Modal <Modal
open={enableModalVisible} open={enableModalVisible}
title={`Enable ${enableTarget ?? ''}`} title={`Enable ${enableTarget ?? ''}`}

View File

@@ -1,4 +1,3 @@
// src/components/Sniffing.tsx
import { import {
CheckCircleOutlined, CheckCircleOutlined,
ExclamationCircleOutlined, ExclamationCircleOutlined,
@@ -23,6 +22,7 @@ import {
Typography, Typography,
} from 'antd'; } from 'antd';
import { ReactElement, useMemo, useState } from 'react'; import { ReactElement, useMemo, useState } from 'react';
import { startSniffer, stopSniffer, stopSnifferByInterface } from '../api/apiClient'; import { startSniffer, stopSniffer, stopSnifferByInterface } from '../api/apiClient';
import { BridgeInfo, InterfaceInfo } from '../types/network'; import { BridgeInfo, InterfaceInfo } from '../types/network';
import { InterfaceSnifferStatus } from '../types/sniffer'; import { InterfaceSnifferStatus } from '../types/sniffer';
@@ -41,69 +41,66 @@ interface SnifferManagerProps {
} }
export default function SnifferManager(props: SnifferManagerProps): ReactElement { export default function SnifferManager(props: SnifferManagerProps): ReactElement {
// local UI state
// modal / form
const [isModalOpen, setIsModalOpen] = useState(false); const [isModalOpen, setIsModalOpen] = useState(false);
const [startMode, setStartMode] = useState<'interface' | 'bridge'>('interface'); const [startMode, setStartMode] = useState<'interface' | 'bridge'>('interface');
const [form] = Form.useForm(); const [form] = Form.useForm();
// derived entries
const statusEntries = useMemo( const statusEntries = useMemo(
() => Object.entries(props.statusMap) as [string, InterfaceSnifferStatus][], () => Object.entries(props.statusMap) as [string, InterfaceSnifferStatus][],
[props.statusMap], [props.statusMap],
); );
// open/close modal
const onOpenStartModal = () => { const onOpenStartModal = () => {
form.resetFields(); form.resetFields();
setStartMode('interface'); setStartMode('interface');
setIsModalOpen(true); setIsModalOpen(true);
}; };
const onCloseModal = () => setIsModalOpen(false);
// start submit const onCloseModal = () => {
const handleStartSubmit = async (values: any) => { setIsModalOpen(false);
const { target } = values; };
const handleStartSubmit = async (values: { target?: string }) => {
const target = values.target;
if (!target) { if (!target) {
notification.warning({ message: 'Warning', description: 'Please select a target to start sniffing on.' }); notification.warning({ message: 'Warning', description: 'Please select a target to start sniffing on.' });
return; return;
} }
try { try {
const payload = startMode === 'interface' ? { interface: target } : { bridge: target }; const payload = startMode === 'interface' ? { interface: target } : { bridge: target };
const res = await startSniffer(payload); const result = await startSniffer(payload);
notification.success({ notification.success({
message: 'Sniffer started', message: 'Sniffer started',
description: `Sniffer started on ${target} (session ${res.session_id})`, description: `Sniffer started on ${target} (session ${result.session_id})`,
}); });
await props.refreshAll(); await props.refreshAll();
setIsModalOpen(false); setIsModalOpen(false);
} catch (err: any) { } catch (error: any) {
console.error('startSniffer error', err); console.error('startSniffer error', error);
notification.error(err?.message ?? 'Failed to start sniffer'); notification.error({
} finally { message: 'Failed to start sniffer',
description: error?.message ?? 'Failed to start sniffer',
});
} }
}; };
// stop all
const handleStopAll = async () => { const handleStopAll = async () => {
try { try {
await stopSniffer(); await stopSniffer();
notification.success({ notification.success({ message: 'All sniffers stopped' });
message: 'All sniffers stopped',
});
await props.refreshStatus(); await props.refreshStatus();
} catch (err: any) { } catch (error: any) {
console.error('stopSniffer error', err); console.error('stopSniffer error', error);
notification.error(err?.message ?? 'Failed to stop sniffers'); notification.error({
} finally { message: 'Failed to stop sniffers',
description: error?.message ?? 'Failed to stop sniffers',
});
} }
}; };
// stop per-interface const handleStopFromList = async (ifaceName: string, sessionId?: string | null) => {
const handleStopFromList = async (ifaceName: string, session_id?: string | null) => {
try { try {
// prefer stop by interface
await stopSnifferByInterface(ifaceName); await stopSnifferByInterface(ifaceName);
notification.success({ notification.success({
message: 'Sniffer stopped', message: 'Sniffer stopped',
@@ -111,28 +108,28 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
}); });
await props.refreshStatus(); await props.refreshStatus();
return; return;
} catch (err: any) { } catch (error: any) {
console.error('stopSnifferByInterface error', err); console.error('stopSnifferByInterface error', error);
// try stop by session as fallback
if (session_id) {
try {
await stopSniffer({ session_id });
notification.success({
message: 'Sniffer stopped',
description: `Sniffer stopped on interface ${ifaceName} (session ${session_id})`,
});
await props.refreshStatus();
return;
} catch (err2: any) {
console.error('stopSniffer by session fallback failed', err2);
}
}
notification.error({
message: 'Failed to stop sniffer',
description: err?.message ?? 'Failed to stop sniffer',
});
} finally {
} }
if (sessionId) {
try {
await stopSniffer({ session_id: sessionId });
notification.success({
message: 'Sniffer stopped',
description: `Sniffer stopped on interface ${ifaceName} (session ${sessionId})`,
});
await props.refreshStatus();
return;
} catch (fallbackError: any) {
console.error('stopSniffer by session fallback failed', fallbackError);
}
}
notification.error({
message: 'Failed to stop sniffer',
description: 'Could not stop sniffer for this interface.',
});
}; };
return ( return (
@@ -145,11 +142,9 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
<Tooltip title="Start a new sniffer session"> <Tooltip title="Start a new sniffer session">
<Button icon={<PlusOutlined />} onClick={onOpenStartModal} /> <Button icon={<PlusOutlined />} onClick={onOpenStartModal} />
</Tooltip> </Tooltip>
<Tooltip title="Refresh status"> <Tooltip title="Refresh status">
<Button icon={<ReloadOutlined />} onClick={() => props.refreshStatus()} /> <Button icon={<ReloadOutlined />} onClick={() => props.refreshStatus()} />
</Tooltip> </Tooltip>
<Tooltip title="Stop all sniffers"> <Tooltip title="Stop all sniffers">
<Button danger icon={<StopOutlined />} onClick={handleStopAll} /> <Button danger icon={<StopOutlined />} onClick={handleStopAll} />
</Tooltip> </Tooltip>
@@ -167,22 +162,19 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
) : ( ) : (
<List <List
dataSource={statusEntries} dataSource={statusEntries}
renderItem={([name, st]: [string, InterfaceSnifferStatus]) => { renderItem={([name, status]: [string, InterfaceSnifferStatus]) => {
const running = st.running; const sessionId = status.session_id ?? null;
const exists = st.exists; const sessionLabel = status.session_label ?? null;
const up = st.up;
const session_id = st.session_id ?? null;
const session_label = st.session_label ?? null;
return ( return (
<List.Item <List.Item
actions={[ actions={[
running ? ( status.running ? (
<Button <Button
key="stop" key="stop"
size="small" size="small"
icon={<StopOutlined />} icon={<StopOutlined />}
onClick={() => handleStopFromList(name, session_id)} onClick={() => handleStopFromList(name, sessionId)}
disabled={props.loading} disabled={props.loading}
> >
Stop Stop
@@ -194,21 +186,19 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
type="primary" type="primary"
icon={<PlayCircleOutlined />} icon={<PlayCircleOutlined />}
onClick={async () => { onClick={async () => {
// start on this interface
try { try {
const res = await startSniffer({ interface: name }); const result = await startSniffer({ interface: name });
notification.success({ notification.success({
message: 'Sniffer started', message: 'Sniffer started',
description: `Sniffer started on ${name} (session ${res.session_id})`, description: `Sniffer started on ${name} (session ${result.session_id})`,
}); });
await props.refreshStatus(); await props.refreshStatus();
} catch (err: any) { } catch (error: any) {
console.error('startSniffer quick', err); console.error('startSniffer quick', error);
notification.error({ notification.error({
message: 'Failed to start sniffer', message: 'Failed to start sniffer',
description: err?.message ?? 'Failed to start sniffer', description: error?.message ?? 'Failed to start sniffer',
}); });
} finally {
} }
}} }}
> >
@@ -221,7 +211,7 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
title={ title={
<Space> <Space>
<Text strong>{name}</Text> <Text strong>{name}</Text>
{running ? ( {status.running ? (
<Tag icon={<CheckCircleOutlined />} color="success"> <Tag icon={<CheckCircleOutlined />} color="success">
running running
</Tag> </Tag>
@@ -230,16 +220,14 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
stopped stopped
</Tag> </Tag>
)} )}
{!status.exists && <Tag color="error">missing</Tag>}
{!exists && <Tag color="error">missing</Tag>} {status.exists && !status.up && <Tag color="warning">down</Tag>}
{exists && !up && <Tag color="warning">down</Tag>} {status.exists && status.up && <Tag color="processing">up</Tag>}
{exists && up && <Tag color="processing">up</Tag>} {sessionId && (
{session_id && (
<Tag> <Tag>
{session_label ? `${session_label}` : 'session'}:{' '} {sessionLabel ?? 'session'}:{' '}
<Text code copyable={{ text: session_id }}> <Text code copyable={{ text: sessionId }}>
{session_id.slice(0, 8)} {sessionId.slice(0, 8)}
</Text> </Text>
</Tag> </Tag>
)} )}
@@ -254,7 +242,6 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
)} )}
</Card> </Card>
{/* Start sniffer modal */}
<Modal <Modal
title="Start sniffer session" title="Start sniffer session"
open={isModalOpen} open={isModalOpen}
@@ -267,8 +254,8 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
<Form.Item label="Mode" name="mode"> <Form.Item label="Mode" name="mode">
<Radio.Group <Radio.Group
value={startMode} value={startMode}
onChange={(e) => { onChange={(event) => {
setStartMode(e.target.value); setStartMode(event.target.value);
form.setFieldsValue({ target: undefined }); form.setFieldsValue({ target: undefined });
}} }}
> >
@@ -291,14 +278,14 @@ export default function SnifferManager(props: SnifferManagerProps): ReactElement
} }
> >
{startMode === 'interface' {startMode === 'interface'
? props.interfaces.map((i) => ( ? props.interfaces.map((iface) => (
<Option key={`if:${i.name}`} value={i.name}> <Option key={`if:${iface.name}`} value={iface.name}>
{i.name} {iface.name}
</Option> </Option>
)) ))
: props.bridges.map((b) => ( : props.bridges.map((bridge) => (
<Option key={`br:${b.ifname}`} value={b.ifname}> <Option key={`br:${bridge.ifname}`} value={bridge.ifname}>
{b.ifname} ({b.members.map((m) => m.ifname).join(', ')}) {bridge.ifname} ({bridge.members.map((member) => member.ifname).join(', ')})
</Option> </Option>
))} ))}
</Select> </Select>

View File

@@ -1,7 +1,7 @@
// src/hooks/useNetwork.ts import { useQuery, useQueryClient } from '@tanstack/react-query';
import { useQuery, useQueryClient } from "@tanstack/react-query"; import { useCallback, useState } from 'react';
import { useCallback, useState } from "react";
import * as api from "../api/apiClient"; import * as api from '../api/apiClient';
import type { import type {
BridgeCreateRequest, BridgeCreateRequest,
BridgeInfo, BridgeInfo,
@@ -9,24 +9,15 @@ import type {
FullState, FullState,
InterfaceInfo, InterfaceInfo,
RouteInfo, RouteInfo,
} from "../types/network"; } from '../types/network';
import { SnifferStatusResponse } from "../types/sniffer"; import { SnifferStatusResponse } from '../types/sniffer';
const TEN_SECONDS = 1000 * 10;
const FIVE_SECONDS = 1000 * 5;
/**
* useNetwork
*
* - queries start disabled (no automatic network calls)
* - calling fetchInterfaces()/fetchBridges()/... will:
* 1) fetch and cache the data right away (queryClient.fetchQuery)
* 2) enable the corresponding useQuery so it becomes "active" and will
* auto-refetch based on the query options
*
* This gives "no initial auto-fetch" but "once fetched, auto-updates".
*/
export function useBackendAPI() { export function useBackendAPI() {
const qc = useQueryClient(); const queryClient = useQueryClient();
// per-query enabled flags (start false => no automatic fetch)
const [interfacesEnabled, setInterfacesEnabled] = useState(false); const [interfacesEnabled, setInterfacesEnabled] = useState(false);
const [linksEnabled, setLinksEnabled] = useState(false); const [linksEnabled, setLinksEnabled] = useState(false);
const [routesEnabled, setRoutesEnabled] = useState(false); const [routesEnabled, setRoutesEnabled] = useState(false);
@@ -34,132 +25,126 @@ export function useBackendAPI() {
const [fullStateEnabled, setFullStateEnabled] = useState(false); const [fullStateEnabled, setFullStateEnabled] = useState(false);
const [snifferStatusEnabled, setSnifferStatusEnabled] = useState(false); const [snifferStatusEnabled, setSnifferStatusEnabled] = useState(false);
// common query options once enabled
const commonOptions = { const commonOptions = {
refetchOnWindowFocus: true, refetchOnWindowFocus: true,
staleTime: 1000 * 10, // 10s staleTime: TEN_SECONDS,
}; };
// Queries (disabled initially)
const interfacesQuery = useQuery<InterfaceInfo[]>({ const interfacesQuery = useQuery<InterfaceInfo[]>({
queryKey: ["interfaces"], queryKey: ['interfaces'],
queryFn: api.fetchInterfaces, queryFn: api.fetchInterfaces,
enabled: interfacesEnabled, enabled: interfacesEnabled,
...commonOptions, ...commonOptions,
}); });
const linksQuery = useQuery<InterfaceInfo[]>({ const linksQuery = useQuery<InterfaceInfo[]>({
queryKey: ["links"], queryKey: ['links'],
queryFn: api.fetchLinks, queryFn: api.fetchLinks,
enabled: linksEnabled, enabled: linksEnabled,
...commonOptions, ...commonOptions,
}); });
const routesQuery = useQuery<RouteInfo[]>({ const routesQuery = useQuery<RouteInfo[]>({
queryKey: ["routes"], queryKey: ['routes'],
queryFn: api.fetchRoutes, queryFn: api.fetchRoutes,
enabled: routesEnabled, enabled: routesEnabled,
...commonOptions, ...commonOptions,
}); });
const bridgesQuery = useQuery<BridgeInfo[]>({ const bridgesQuery = useQuery<BridgeInfo[]>({
queryKey: ["bridges"], queryKey: ['bridges'],
queryFn: api.fetchBridges, queryFn: api.fetchBridges,
enabled: bridgesEnabled, enabled: bridgesEnabled,
...commonOptions, ...commonOptions,
}); });
const fullStateQuery = useQuery<FullState>({ const fullStateQuery = useQuery<FullState>({
queryKey: ["full-state"], queryKey: ['full-state'],
queryFn: api.fetchFullState, queryFn: api.fetchFullState,
enabled: fullStateEnabled, enabled: fullStateEnabled,
...commonOptions, ...commonOptions,
}); });
const snifferStatusQuery = useQuery<SnifferStatusResponse>({ const snifferStatusQuery = useQuery<SnifferStatusResponse>({
queryKey: ["sniffer-status"], queryKey: ['sniffer-status'],
queryFn: api.fetchSnifferStatus, queryFn: api.fetchSnifferStatus,
enabled: fullStateEnabled, enabled: snifferStatusEnabled,
...commonOptions, ...commonOptions,
}); });
// Imperative fetch helpers that also enable auto-refetch behavior
const fetchInterfaces = useCallback(async () => { const fetchInterfaces = useCallback(async () => {
const res = await qc.fetchQuery<InterfaceInfo[]>({ const result = await queryClient.fetchQuery<InterfaceInfo[]>({
queryKey: ["interfaces"], queryKey: ['interfaces'],
queryFn: api.fetchInterfaces, queryFn: api.fetchInterfaces,
staleTime: 1000 * 10 staleTime: TEN_SECONDS,
}); });
setInterfacesEnabled(true); setInterfacesEnabled(true);
return res; return result;
}, [qc]); }, [queryClient]);
const fetchLinks = useCallback(async () => { const fetchLinks = useCallback(async () => {
const res = await qc.fetchQuery<InterfaceInfo[]>({ const result = await queryClient.fetchQuery<InterfaceInfo[]>({
queryKey: ["links"], queryKey: ['links'],
queryFn: api.fetchLinks, queryFn: api.fetchLinks,
staleTime: 1000 * 10 staleTime: TEN_SECONDS,
}); });
setLinksEnabled(true); setLinksEnabled(true);
return res; return result;
}, [qc]); }, [queryClient]);
const fetchRoutes = useCallback(async () => { const fetchRoutes = useCallback(async () => {
const res = await qc.fetchQuery<RouteInfo[]>({ const result = await queryClient.fetchQuery<RouteInfo[]>({
queryKey: ["routes"], queryKey: ['routes'],
queryFn: api.fetchRoutes, queryFn: api.fetchRoutes,
staleTime: 1000 * 10 staleTime: TEN_SECONDS,
}); });
setRoutesEnabled(true); setRoutesEnabled(true);
return res; return result;
}, [qc]); }, [queryClient]);
const fetchBridges = useCallback(async () => { const fetchBridges = useCallback(async () => {
const res = await qc.fetchQuery<BridgeInfo[]>({ const result = await queryClient.fetchQuery<BridgeInfo[]>({
queryKey: ["bridges"], queryKey: ['bridges'],
queryFn: api.fetchBridges, queryFn: api.fetchBridges,
staleTime: 1000 * 10 staleTime: TEN_SECONDS,
}); });
setBridgesEnabled(true); setBridgesEnabled(true);
return res; return result;
}, [qc]); }, [queryClient]);
const fetchFullState = useCallback(async () => { const fetchFullState = useCallback(async () => {
const res = await qc.fetchQuery<FullState>({ const result = await queryClient.fetchQuery<FullState>({
queryKey: ["full-state"], queryKey: ['full-state'],
queryFn: api.fetchFullState, queryFn: api.fetchFullState,
staleTime: 1000 * 5 staleTime: FIVE_SECONDS,
}); });
setFullStateEnabled(true); setFullStateEnabled(true);
return res; return result;
}, [qc]); }, [queryClient]);
const fetchSnifferStatus = useCallback(async () => { const fetchSnifferStatus = useCallback(async () => {
const res = await qc.fetchQuery<SnifferStatusResponse>({ const result = await queryClient.fetchQuery<SnifferStatusResponse>({
queryKey: ["sniffer-status"], queryKey: ['sniffer-status'],
queryFn: api.fetchSnifferStatus, queryFn: api.fetchSnifferStatus,
staleTime: 1000 * 5 staleTime: FIVE_SECONDS,
}); });
setFullStateEnabled(true); setSnifferStatusEnabled(true);
return res; return result;
}, [qc]); }, [queryClient]);
// Local loading state for simple UI feedback
const [isCreating, setIsCreating] = useState(false); const [isCreating, setIsCreating] = useState(false);
const [isRemoving, setIsRemoving] = useState(false); const [isRemoving, setIsRemoving] = useState(false);
// Simple imperative functions that call the API and invalidate queries
async function createBridge(payload: BridgeCreateRequest) { async function createBridge(payload: BridgeCreateRequest) {
setIsCreating(true); setIsCreating(true);
try { try {
await api.createBridge(payload); await api.createBridge(payload);
// If the query is enabled it will refetch automatically after invalidation. await queryClient.invalidateQueries({ queryKey: ['bridges'] });
await qc.invalidateQueries({ queryKey: ["bridges"] }); await queryClient.invalidateQueries({ queryKey: ['full-state'] });
await qc.invalidateQueries({ queryKey: ["full-state"] }); await queryClient.invalidateQueries({ queryKey: ['interfaces'] });
await qc.invalidateQueries({ queryKey: ["interfaces"] }); } catch (error) {
} catch (err) {
const message = const message =
err instanceof Error ? err.message : typeof err === "string" ? err : "Create bridge failed"; error instanceof Error ? error.message : typeof error === 'string' ? error : 'Create bridge failed';
throw new Error(message); throw new Error(message);
} finally { } finally {
setIsCreating(false); setIsCreating(false);
@@ -170,45 +155,43 @@ export function useBackendAPI() {
setIsRemoving(true); setIsRemoving(true);
try { try {
await api.removeBridge(payload); await api.removeBridge(payload);
await qc.invalidateQueries({ queryKey: ["bridges"] }); await queryClient.invalidateQueries({ queryKey: ['bridges'] });
await qc.invalidateQueries({ queryKey: ["full-state"] }); await queryClient.invalidateQueries({ queryKey: ['full-state'] });
await qc.invalidateQueries({ queryKey: ["interfaces"] }); await queryClient.invalidateQueries({ queryKey: ['interfaces'] });
await qc.invalidateQueries({ queryKey: ["sniffer-status"] }); await queryClient.invalidateQueries({ queryKey: ['sniffer-status'] });
} catch (err) { } catch (error) {
const message = const message =
err instanceof Error ? err.message : typeof err === "string" ? err : "Remove bridge failed"; error instanceof Error ? error.message : typeof error === 'string' ? error : 'Remove bridge failed';
throw new Error(message); throw new Error(message);
} finally { } finally {
setIsRemoving(false); setIsRemoving(false);
} }
} }
// Convenience: invalidate helpers
function refreshInterfaces() { function refreshInterfaces() {
return qc.invalidateQueries({ queryKey: ["interfaces"] }); return queryClient.invalidateQueries({ queryKey: ['interfaces'] });
} }
function refreshLinks() { function refreshLinks() {
return qc.invalidateQueries({ queryKey: ["links"] }); return queryClient.invalidateQueries({ queryKey: ['links'] });
} }
function refreshRoutes() { function refreshRoutes() {
return qc.invalidateQueries({ queryKey: ["routes"] }); return queryClient.invalidateQueries({ queryKey: ['routes'] });
} }
function refreshBridges() { function refreshBridges() {
return qc.invalidateQueries({ queryKey: ["bridges"] }); return queryClient.invalidateQueries({ queryKey: ['bridges'] });
} }
function refreshFullState() { function refreshFullState() {
return qc.invalidateQueries({ queryKey: ["full-state"] }); return queryClient.invalidateQueries({ queryKey: ['full-state'] });
} }
function refreshSnifferStatus() { function refreshSnifferStatus() {
return qc.invalidateQueries({ queryKey: ["sniffer-status"] }); return queryClient.invalidateQueries({ queryKey: ['sniffer-status'] });
} }
// Convenience: refresh all queries
function refreshAll() { function refreshAll() {
refreshInterfaces(); refreshInterfaces();
refreshLinks(); refreshLinks();
@@ -219,39 +202,28 @@ export function useBackendAPI() {
} }
return { return {
// queries
interfacesQuery, interfacesQuery,
linksQuery, linksQuery,
routesQuery, routesQuery,
bridgesQuery, bridgesQuery,
fullStateQuery, fullStateQuery,
snifferStatusQuery, snifferStatusQuery,
// manual fetchers (fetch+enable auto-updates)
fetchInterfaces, fetchInterfaces,
fetchLinks, fetchLinks,
fetchRoutes, fetchRoutes,
fetchBridges, fetchBridges,
fetchFullState, fetchFullState,
fetchSnifferStatus, fetchSnifferStatus,
// simple mutation functions (imperative)
createBridge, createBridge,
removeBridge, removeBridge,
// local loading flags
isCreating, isCreating,
isRemoving, isRemoving,
// invalidate helpers
refreshInterfaces, refreshInterfaces,
refreshLinks, refreshLinks,
refreshRoutes, refreshRoutes,
refreshBridges, refreshBridges,
refreshFullState, refreshFullState,
refreshSnifferStatus, refreshSnifferStatus,
// refresh all
refreshAll, refreshAll,
}; };
} }

View File

@@ -1,4 +1,3 @@
// src/pages/BridgesManager.tsx
import { DeleteOutlined, PlusOutlined, ReloadOutlined } from '@ant-design/icons'; import { DeleteOutlined, PlusOutlined, ReloadOutlined } from '@ant-design/icons';
import { import {
Button, Button,
@@ -18,9 +17,9 @@ import {
} from 'antd'; } from 'antd';
import type { ColumnsType } from 'antd/es/table'; import type { ColumnsType } from 'antd/es/table';
import { useEffect, useMemo, useState } from 'react'; import { useEffect, useMemo, useState } from 'react';
import { createBridge, fetchFullState, removeBridge } from '../api/apiClient'; import { createBridge, fetchFullState, removeBridge } from '../api/apiClient';
import type { BridgeInfo, InterfaceInfo } from '../types/network'; import type { BridgeInfo, FullState, InterfaceInfo } from '../types/network';
import { FullState } from '../types/network';
const { Title, Paragraph } = Typography; const { Title, Paragraph } = Typography;
@@ -29,19 +28,19 @@ export default function Network() {
const [bridgeForm] = Form.useForm(); const [bridgeForm] = Form.useForm();
const [networkState, setNetworkState] = useState<FullState>(); const [networkState, setNetworkState] = useState<FullState>();
const getFullState = (auto: boolean = false) => { const getFullState = (silent = false) => {
fetchFullState() fetchFullState()
.then((interfaces) => { .then((state) => {
setNetworkState(interfaces); setNetworkState(state);
if (!auto) { if (!silent) {
notification.success({ notification.success({
message: 'Success', message: 'Success',
description: 'Network state updated.', description: 'Network state updated.',
}); });
} }
}) })
.catch((err) => { .catch((error) => {
console.error('Failed to fetch interfaces:', err); console.error('Failed to fetch network state:', error);
notification.error({ notification.error({
message: 'Error', message: 'Error',
description: 'Failed to fetch network state.', description: 'Failed to fetch network state.',
@@ -49,39 +48,36 @@ export default function Network() {
}); });
}; };
// fetch on mount (explicit, since queries are disabled by default)
useEffect(() => { useEffect(() => {
getFullState(true); getFullState(true);
}, []); }, []);
// build select options from interfaces list
const interfaceOptions = useMemo( const interfaceOptions = useMemo(
() => () =>
networkState?.interfaces.map((it: InterfaceInfo) => ({ networkState?.interfaces.map((iface: InterfaceInfo) => ({
label: it.name, label: iface.name,
value: it.name, value: iface.name,
})) ?? [], })) ?? [],
[networkState], [networkState],
); );
// Columns for interfaces table (read-only)
const interfaceColumns: ColumnsType<InterfaceInfo> = useMemo( const interfaceColumns: ColumnsType<InterfaceInfo> = useMemo(
() => [ () => [
{ title: 'IfIndex', dataIndex: 'ifindex', key: 'ifindex', width: 90 }, { title: 'IfIndex', dataIndex: 'ifindex', key: 'ifindex', width: 90 },
{ title: 'Name', dataIndex: 'name', key: 'name' }, { title: 'Name', dataIndex: 'name', key: 'name' },
{ title: 'State', dataIndex: 'state', key: 'state', render: (s) => <Tag>{s}</Tag> }, { title: 'State', dataIndex: 'state', key: 'state', render: (state) => <Tag>{state}</Tag> },
{ title: 'MAC', dataIndex: 'mac', key: 'mac', render: (m) => m ?? '—' }, { title: 'MAC', dataIndex: 'mac', key: 'mac', render: (mac) => mac ?? '—' },
{ title: 'MTU', dataIndex: 'mtu', key: 'mtu', width: 90 }, { title: 'MTU', dataIndex: 'mtu', key: 'mtu', width: 90 },
{ {
title: 'Addresses', title: 'Addresses',
dataIndex: 'addresses', dataIndex: 'addresses',
key: 'addresses', key: 'addresses',
render: (addrs: any[]) => render: (addresses: any[]) =>
addrs?.length ? ( addresses?.length ? (
<Space orientation="vertical"> <Space direction="vertical">
{addrs.map((a) => ( {addresses.map((address) => (
<span key={`${a.address}/${a.prefixlen}`}> <span key={`${address.address}/${address.prefixlen}`}>
{a.address}/{a.prefixlen} ({a.family}) {address.address}/{address.prefixlen} ({address.family})
</span> </span>
))} ))}
</Space> </Space>
@@ -93,23 +89,22 @@ export default function Network() {
[], [],
); );
// Columns for bridges table (with remove action)
const bridgeColumns: ColumnsType<BridgeInfo> = useMemo( const bridgeColumns: ColumnsType<BridgeInfo> = useMemo(
() => [ () => [
{ title: 'IfIndex', dataIndex: 'ifindex', key: 'ifindex', width: 90 }, { title: 'IfIndex', dataIndex: 'ifindex', key: 'ifindex', width: 90 },
{ title: 'Name', dataIndex: 'ifname', key: 'ifname' }, { title: 'Name', dataIndex: 'ifname', key: 'ifname' },
{ title: 'State', dataIndex: 'state', key: 'state', render: (s) => <Tag>{s ?? '—'}</Tag> }, { title: 'State', dataIndex: 'state', key: 'state', render: (state) => <Tag>{state ?? '—'}</Tag> },
{ {
title: 'Members', title: 'Members',
dataIndex: 'members', dataIndex: 'members',
key: 'members', key: 'members',
render: (members: any[]) => (members?.length ? members.map((m) => m.ifname).join(', ') : '—'), render: (members: any[]) => (members?.length ? members.map((member) => member.ifname).join(', ') : '—'),
}, },
{ {
title: 'Actions', title: 'Actions',
key: 'actions', key: 'actions',
width: 140, width: 140,
render: (_: any, record: BridgeInfo) => ( render: (_, record: BridgeInfo) => (
<Popconfirm <Popconfirm
title={`Remove bridge ${record.ifname}?`} title={`Remove bridge ${record.ifname}?`}
onConfirm={() => handleRemoveBridge(record.ifname)} onConfirm={() => handleRemoveBridge(record.ifname)}
@@ -121,11 +116,12 @@ export default function Network() {
), ),
}, },
], ],
[networkState], [],
); );
async function handleCreateBridge(values: { name: string; interfaces?: string[] }) {
const ifaceList = values.interfaces ?? []; function handleCreateBridge(values: { name: string; interfaces?: string[] }) {
createBridge({ name: values.name, interfaces: ifaceList }) const interfaces = values.interfaces ?? [];
createBridge({ name: values.name, interfaces })
.then(() => { .then(() => {
notification.success({ notification.success({
message: 'Success', message: 'Success',
@@ -135,11 +131,11 @@ export default function Network() {
getFullState(true); getFullState(true);
bridgeForm.resetFields(); bridgeForm.resetFields();
}) })
.catch((err) => { .catch((error) => {
console.error(err); console.error(error);
notification.error({ notification.error({
message: 'Error', message: 'Error',
description: (err as Error).message ?? 'Failed to create bridge', description: (error as Error).message ?? 'Failed to create bridge',
}); });
}); });
} }
@@ -153,11 +149,11 @@ export default function Network() {
}); });
getFullState(true); getFullState(true);
}) })
.catch((err) => { .catch((error) => {
console.error(err); console.error(error);
notification.error({ notification.error({
message: 'Error', message: 'Error',
description: (err as Error).message ?? 'Failed to remove bridge', description: (error as Error).message ?? 'Failed to remove bridge',
}); });
}); });
} }
@@ -166,28 +162,21 @@ export default function Network() {
<div style={{ padding: 16 }}> <div style={{ padding: 16 }}>
<Row justify="space-between" align="middle" style={{ marginBottom: 12 }}> <Row justify="space-between" align="middle" style={{ marginBottom: 12 }}>
<Col> <Col>
<Title level={2}> Network Management</Title> <Title level={2}>Network Management</Title>
<Paragraph type="secondary">View system interfaces and manage network bridges.</Paragraph> <Paragraph type="secondary">View system interfaces and manage network bridges.</Paragraph>
</Col> </Col>
<Col> <Col>
<Space> <Button icon={<ReloadOutlined />} onClick={() => getFullState()}>
<Button Refresh
icon={<ReloadOutlined />} </Button>
onClick={() => {
getFullState();
}}
>
Refresh
</Button>
</Space>
</Col> </Col>
</Row> </Row>
<Row gutter={16}> <Row gutter={16}>
<Col span={14}> <Col span={14}>
<Card title={`Interfaces (${networkState?.interfaces.length})`} style={{ overflow: 'auto' }}> <Card title={`Interfaces (${networkState?.interfaces.length ?? 0})`} style={{ overflow: 'auto' }}>
<Table <Table
rowKey={(r: InterfaceInfo) => r.ifindex} rowKey={(row: InterfaceInfo) => row.ifindex}
dataSource={networkState?.interfaces ?? []} dataSource={networkState?.interfaces ?? []}
columns={interfaceColumns} columns={interfaceColumns}
pagination={{ pageSize: 8 }} pagination={{ pageSize: 8 }}
@@ -198,14 +187,13 @@ export default function Network() {
<Col span={10}> <Col span={10}>
<Card <Card
title={`Bridges (${networkState?.bridges.length})`} title={`Bridges (${networkState?.bridges.length ?? 0})`}
style={{ overflow: 'auto' }} style={{ overflow: 'auto' }}
extra={ extra={
<Button <Button
icon={<PlusOutlined />} icon={<PlusOutlined />}
type="primary" type="primary"
onClick={() => { onClick={() => {
// ensure up-to-date interface list when opening modal
getFullState(true); getFullState(true);
setBridgeModalVisible(true); setBridgeModalVisible(true);
}} }}
@@ -213,7 +201,7 @@ export default function Network() {
} }
> >
<Table <Table
rowKey={(r: BridgeInfo) => String(r.ifindex)} rowKey={(row: BridgeInfo) => String(row.ifindex)}
dataSource={networkState?.bridges ?? []} dataSource={networkState?.bridges ?? []}
columns={bridgeColumns} columns={bridgeColumns}
pagination={{ pageSize: 6 }} pagination={{ pageSize: 6 }}
@@ -223,7 +211,6 @@ export default function Network() {
</Col> </Col>
</Row> </Row>
{/* Create Bridge Modal */}
<Modal <Modal
title="Create Bridge" title="Create Bridge"
open={bridgeModalVisible} open={bridgeModalVisible}
@@ -238,7 +225,6 @@ export default function Network() {
<Input placeholder="e.g. br0" /> <Input placeholder="e.g. br0" />
</Form.Item> </Form.Item>
{/* Select field populated from interfaces endpoint */}
<Form.Item name="interfaces" label="Interfaces (select one or more)"> <Form.Item name="interfaces" label="Interfaces (select one or more)">
<Select <Select
mode="multiple" mode="multiple"

View File

@@ -1,6 +1,6 @@
// src/components/Sniffing.tsx import { Col, message, Row, Typography } from 'antd';
import { Col, message, Row, Select, Typography } from 'antd';
import { ReactElement, useCallback, useEffect, useState } from 'react'; import { ReactElement, useCallback, useEffect, useState } from 'react';
import { fetchBridges, fetchInterfaces, fetchSnifferStatus } from '../api/apiClient'; import { fetchBridges, fetchInterfaces, fetchSnifferStatus } from '../api/apiClient';
import PacketViewer from '../components/PacketViewer'; import PacketViewer from '../components/PacketViewer';
import SnifferManager from '../components/SnifferManager'; import SnifferManager from '../components/SnifferManager';
@@ -8,48 +8,48 @@ import { BridgeInfo, InterfaceInfo } from '../types/network';
import { InterfaceSnifferStatus } from '../types/sniffer'; import { InterfaceSnifferStatus } from '../types/sniffer';
const { Title, Text } = Typography; const { Title, Text } = Typography;
const { Option } = Select;
export default function Sniffing(): ReactElement { export default function Sniffing(): ReactElement {
const [interfaces, setInterfaces] = useState<InterfaceInfo[]>([]); const [interfaces, setInterfaces] = useState<InterfaceInfo[]>([]);
const [bridges, setBridges] = useState<BridgeInfo[]>([]); const [bridges, setBridges] = useState<BridgeInfo[]>([]);
const [statusMap, setStatusMap] = useState<Record<string, InterfaceSnifferStatus>>({}); const [statusMap, setStatusMap] = useState<Record<string, InterfaceSnifferStatus>>({});
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
const [statusLoading, setStatusLoading] = useState(false); const [statusLoading, setStatusLoading] = useState(false);
// initial load const refreshStatus = useCallback(async () => {
useEffect(() => { setStatusLoading(true);
refreshAll().catch(() => {}); try {
const status = await fetchSnifferStatus();
setStatusMap(status.interfaces ?? {});
} catch (error: any) {
console.error('fetchSnifferStatus error', error);
message.error(error?.message ?? 'Failed to fetch sniffer status');
} finally {
setStatusLoading(false);
}
}, []); }, []);
const refreshAll = useCallback(async () => { const refreshAll = useCallback(async () => {
setLoading(true); setLoading(true);
try { try {
const [ifs, brs] = await Promise.allSettled([fetchInterfaces(), fetchBridges()]); const [ifaces, bridgesResult] = await Promise.allSettled([fetchInterfaces(), fetchBridges()]);
if (ifs.status === 'fulfilled') setInterfaces(ifs.value); if (ifaces.status === 'fulfilled') {
if (brs.status === 'fulfilled') setBridges(brs.value); setInterfaces(ifaces.value);
}
if (bridgesResult.status === 'fulfilled') {
setBridges(bridgesResult.value);
}
await refreshStatus(); await refreshStatus();
} catch (err) {
// ignore; errors handled in individual calls
} finally { } finally {
setLoading(false); setLoading(false);
} }
}, []); }, [refreshStatus]);
useEffect(() => {
refreshAll().catch(() => undefined);
}, [refreshAll]);
const refreshStatus = useCallback(async () => {
setStatusLoading(true);
try {
const st = await fetchSnifferStatus();
setStatusMap(st.interfaces ?? {});
} catch (err: any) {
console.error('fetchSnifferStatus error', err);
message.error(err?.message ?? 'Failed to fetch sniffer status');
} finally {
setStatusLoading(false);
}
}, []);
return ( return (
<div className="sniffing-page" style={{ padding: 16 }}> <div className="sniffing-page" style={{ padding: 16 }}>
<Row justify="space-between" align="middle" style={{ marginBottom: 12 }}> <Row justify="space-between" align="middle" style={{ marginBottom: 12 }}>
@@ -57,9 +57,10 @@ export default function Sniffing(): ReactElement {
<Title level={2} style={{ margin: 0 }}> <Title level={2} style={{ margin: 0 }}>
Sniffing Sniffing
</Title> </Title>
<Text type="secondary">Start, stop and view AF_PACKET sniffer sessions</Text> <Text type="secondary">Start, stop and view AF_PACKET sniffer sessions.</Text>
</Col> </Col>
</Row> </Row>
<Row> <Row>
<Col span={24}> <Col span={24}>
<SnifferManager <SnifferManager
@@ -73,9 +74,10 @@ export default function Sniffing(): ReactElement {
/> />
</Col> </Col>
</Row> </Row>
<Row> <Row>
<Col span={24} style={{ marginTop: 24 }}> <Col span={24} style={{ marginTop: 24 }}>
<PacketViewer interfaces={interfaces.map((i) => i.name)} /> <PacketViewer interfaces={interfaces.map((iface) => iface.name)} />
</Col> </Col>
</Row> </Row>
</div> </div>