From b60a1d311827feee7a6deaf359a03ff3a7f3f3af Mon Sep 17 00:00:00 2001 From: malmert Date: Fri, 6 Mar 2026 19:03:26 +0100 Subject: [PATCH] file structure and comments unified --- backend/src/Models/netplan.py | 42 +- backend/src/Models/packets.py | 16 +- backend/src/api/network_api.py | 361 ++++----- backend/src/api/packet_api.py | 170 ++--- backend/src/api/sniffer_api.py | 184 ++--- backend/src/main.py | 77 +- backend/src/shared_objects.py | 12 +- backend/src/utilities/database.py | 91 +-- backend/src/utilities/packet_broadcaster.py | 104 +-- frontend/src/api/apiClient.ts | 127 +--- frontend/src/appRouter.tsx | 29 +- .../src/components/FireWallAddChainModal.tsx | 19 +- .../src/components/FireWallAddTableModal.tsx | 15 +- .../src/components/FirewallRuleBuilder.tsx | 718 +----------------- .../src/components/FirewallRulesetViewer.tsx | 33 +- frontend/src/components/ScriptManager.tsx | 55 +- frontend/src/components/SnifferManager.tsx | 155 ++-- frontend/src/hooks/useBackendAPI.ts | 160 ++-- frontend/src/pages/Network.tsx | 96 +-- frontend/src/pages/Sniffing.tsx | 56 +- 20 files changed, 697 insertions(+), 1823 deletions(-) diff --git a/backend/src/Models/netplan.py b/backend/src/Models/netplan.py index 625936d..8856fa8 100644 --- a/backend/src/Models/netplan.py +++ b/backend/src/Models/netplan.py @@ -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 typing import List, Dict, Optional class Nameservers(BaseModel): + """DNS nameserver configuration.""" + addresses: List[str] = Field(default_factory=list) search: List[str] = Field(default_factory=list) class EthernetConfig(BaseModel): + """Netplan ethernet interface configuration.""" + dhcp4: Optional[bool] = None dhcp6: Optional[bool] = None addresses: Optional[List[str]] = None @@ -18,44 +25,23 @@ class EthernetConfig(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 dhcp6: Optional[bool] = None addresses: Optional[List[str]] = None gateway4: Optional[str] = None gateway6: Optional[str] = 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 class NetworkConfig(BaseModel): + """Top-level Netplan network object.""" + version: int = 2 renderer: Optional[str] = "networkd" ethernets: Dict[str, EthernetConfig] = 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 - } - } -}''' diff --git a/backend/src/Models/packets.py b/backend/src/Models/packets.py index a201797..a19fa37 100644 --- a/backend/src/Models/packets.py +++ b/backend/src/Models/packets.py @@ -1,12 +1,16 @@ -# src/models/packet.py -from datetime import datetime +"""Pydantic model for packet rows returned by the backend.""" -from pydantic import BaseModel, Field, IPvAnyAddress +from datetime import datetime from typing import Optional, Union +from pydantic import BaseModel, Field, IPvAnyAddress + + class PacketDBModel(BaseModel): + """Normalized packet representation used across DB and API layers.""" + id: Union[int, str] - timestamp: datetime = Field(..., description="ISO timestamp") + timestamp: datetime = Field(..., description="Packet timestamp in ISO format.") iface: str src_mac: Optional[str] = None dst_mac: Optional[str] = None @@ -18,7 +22,7 @@ class PacketDBModel(BaseModel): dst_port: Optional[int] = None vlan_id: 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 packets: Optional[int] = None @@ -40,4 +44,4 @@ class PacketDBModel(BaseModel): "length": 128, "raw_b64": "BASE64...", } - } \ No newline at end of file + } diff --git a/backend/src/api/network_api.py b/backend/src/api/network_api.py index 2ae3c29..120831e 100644 --- a/backend/src/api/network_api.py +++ b/backend/src/api/network_api.py @@ -1,122 +1,95 @@ +"""Network inspection and bridge management endpoints.""" + +from typing import List, Optional + from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel, Field -from typing import List, Optional from pyroute2 import IPRoute, NDB router = APIRouter() -# Globals for lazy initialization ip: IPRoute | None = None ndb: NDB | None = None -# ------------------------------ -# Pydantic models -# ------------------------------ class InterfaceAddress(BaseModel): - """ - Represents an IP address assigned to a network interface. - """ + """IP address assigned to an interface.""" + family: str = Field(..., description="IP family: 'ipv4' or 'ipv6'.") - address: str = Field(..., description="The IP address assigned to the interface.") - prefixlen: int = Field(..., description="Subnet prefix length (e.g., 24 for 255.255.255.0).") + address: str = Field(..., description="IP address.") + prefixlen: int = Field(..., description="Subnet prefix length.") class InterfaceInfo(BaseModel): - """ - Represents a network interface with all its properties. - """ - ifindex: int = Field(..., description="Interface index (unique identifier assigned by the kernel).") - name: str = Field(..., description="Interface name (e.g., 'eth0', 'enp38s0').") - state: str = Field(..., description="Operational state (e.g., 'UP', 'DOWN', 'UNKNOWN').") - mac: Optional[str] = Field(None, description="MAC address of the interface, if applicable.") - mtu: int = Field(..., description="Maximum Transmission Unit for the interface.") - flags: List[str] = Field(..., description="List of interface flags (e.g., ['BROADCAST', 'MULTICAST']).") - addresses: List[InterfaceAddress] = Field(..., description="List of IP addresses assigned to the interface.") + """Interface with link metadata and assigned addresses.""" + + ifindex: int = Field(..., description="Kernel interface index.") + name: str = Field(..., description="Interface name.") + state: str = Field(..., description="Operational state.") + mac: Optional[str] = Field(None, description="MAC address.") + mtu: int = Field(..., description="Maximum transmission unit.") + flags: List[str] = Field(..., description="Decoded interface flags.") + addresses: List[InterfaceAddress] = Field(..., description="Assigned IP addresses.") + class RouteInfo(BaseModel): - """ - Represents a single routing table entry. - """ + """Single routing table entry.""" - dst: Optional[str] = Field( - None, description="Destination network in CIDR notation (e.g., '192.168.1.0/24'). None means default route." - ) + dst: Optional[str] = Field(None, description="Destination CIDR; null 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): - """ - Represents a network interface which is a member of an bridge. - """ - ifindex: int = Field(..., description="Interface index of a bridge member") - ifname: str = Field(..., description="Interface name of a bridge member") - state: Optional[str] = Field(None, description="Operational state of the interface") - mtu: Optional[int] = Field(None, description="MTU of the interface") + """Interface that belongs to a bridge.""" + + ifindex: int = Field(..., description="Interface index.") + ifname: str = Field(..., description="Interface name.") + state: Optional[str] = Field(None, description="Operational state.") + mtu: Optional[int] = Field(None, description="Interface MTU.") + class BridgeInfo(BaseModel): - """ - Represents a network bridge interface with all its properties. - """ - ifindex: int = Field(..., description="Interface index of the bridge") - ifname: str = Field(..., description="Bridge interface name") - state: Optional[str] = Field(None, description="Operational state of the bridge") - mtu: Optional[int] = Field(None, description="MTU of the bridge") - stp_state: Optional[int] = Field(None, description="STP (Spanning Tree Protocol) state of the bridge") - members: List[BridgeInterfaceInfo] = Field(default_factory=list, description="List of member interfaces of the bridge") + """Bridge interface with member information.""" + + ifindex: int = Field(..., description="Bridge index.") + ifname: str = Field(..., description="Bridge name.") + state: Optional[str] = Field(None, description="Bridge state.") + mtu: Optional[int] = Field(None, description="Bridge MTU.") + stp_state: Optional[int] = Field(None, description="Spanning tree state.") + members: List[BridgeInterfaceInfo] = Field(default_factory=list, description="Bridge members.") + class BridgeCreateRequest(BaseModel): + """Payload for creating a bridge and attaching interfaces.""" + name: str interfaces: List[str] class BridgeRemoveRequest(BaseModel): - name: str -# ------------------------------ -# Lazy Init Functions -# ------------------------------ + """Payload for removing a bridge.""" -def init_network_api(): + name: str + + +def init_network_api() -> None: + """Initialize lazy pyroute2 clients.""" global ip, ndb if ip is None: ip = IPRoute() if ndb is None: ndb = NDB() -def shutdown_network_api(): + +def shutdown_network_api() -> None: + """Close pyroute2 clients if they were initialized.""" global ip, ndb if ip: ip.close() @@ -125,37 +98,38 @@ def shutdown_network_api(): ndb.close() ndb = None -def get_iproute(): + +def get_iproute() -> IPRoute: + """Dependency provider for the shared IPRoute instance.""" if ip is None: init_network_api() return ip -def get_ndb(): + +def get_ndb() -> NDB: + """Dependency provider for the shared NDB instance.""" if ndb is None: init_network_api() return ndb -# ------------------------------ -# Utility functions -# ------------------------------ -def parse_addresses(addrs): - res = [] - for a in addrs: - family = "ipv4" if a.get("family") == 2 else "ipv6" - res.append( +def parse_addresses(addrs: list[dict]) -> list[InterfaceAddress]: + """Convert pyroute2 address rows into `InterfaceAddress` models.""" + result: list[InterfaceAddress] = [] + for addr in addrs: + family = "ipv4" if addr.get("family") == 2 else "ipv6" + result.append( InterfaceAddress( family=family, - address=a.get("address"), - prefixlen=a.get("prefixlen"), + address=addr.get("address"), + prefixlen=addr.get("prefixlen"), ) ) - return res + return result + def parse_flags(flags_int: int) -> list[str]: - """ - Converts the integer flags from pyroute2 to human-readable list of strings. - """ + """Decode Linux interface flag bitset to names.""" flags_map = { 0x1: "UP", 0x2: "BROADCAST", @@ -177,37 +151,33 @@ def parse_flags(flags_int: int) -> list[str]: 0x20000: "DORMANT", 0x40000: "ECHO", } - result = [] - for bit, name in flags_map.items(): - if flags_int & bit: - result.append(name) - return result + return [name for bit, name in flags_map.items() if flags_int & bit] -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: raise HTTPException(status_code=404, detail=f"Interface {name} not found") return idx[0] -def bridge_exists(name: str, ip: IPRoute) -> bool: - return bool(ip.link_lookup(ifname=name)) +def bridge_exists(name: str, ip_route: IPRoute) -> bool: + """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]) -def get_interfaces(ip: IPRoute = Depends(get_iproute)): - result = [] +def get_interfaces(ip: IPRoute = Depends(get_iproute)) -> List[InterfaceInfo]: + """List host interfaces with addresses and decoded flags.""" + result: list[InterfaceInfo] = [] links = ip.get_links() addresses = ip.get_addr() - addr_map = {} - for a in addresses: - ifindex = a.get("index") - addr_map.setdefault(ifindex, []).append(a) + addr_map: dict[int, list] = {} + for addr in addresses: + ifindex = addr.get("index") + addr_map.setdefault(ifindex, []).append(addr) for link in links: attrs = dict(link["attrs"]) @@ -227,55 +197,53 @@ def get_interfaces(ip: IPRoute = Depends(get_iproute)): ) return result + @router.get("/routes", response_model=List[RouteInfo]) -def get_routes(ip: IPRoute = Depends(get_iproute)): - routes = [] - for r in ip.get_routes(): - attrs = dict(r["attrs"]) +def get_routes(ip: IPRoute = Depends(get_iproute)) -> List[RouteInfo]: + """List routes from the kernel routing tables.""" + routes: list[RouteInfo] = [] + for route in ip.get_routes(): + attrs = dict(route["attrs"]) dst = attrs.get("RTA_DST") gateway = attrs.get("RTA_GATEWAY") prefsrc = attrs.get("RTA_PREFSRC") - oif = r.get("oif") + oif = route.get("oif") + ifname = None if oif is not None: - # translate ifindex → name link = ip.get_links(oif)[0] ifname = dict(link["attrs"]).get("IFLA_IFNAME") routes.append( 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, prefsrc=prefsrc, oif=oif, ifname=ifname, - table=r.get("table", 254), - proto=r.get("proto"), - scope=r.get("scope"), - type=r.get("type"), + table=route.get("table", 254), + proto=route.get("proto"), + scope=route.get("scope"), + type=route.get("type"), ) ) return routes @router.get("/links", response_model=List[InterfaceInfo]) -def get_raw_links(ip: IPRoute = Depends(get_iproute)): - """ - Returns all interfaces in a clean Pydantic format. - This is similar to /interfaces but avoids additional processing if needed. - """ - result = [] +def get_raw_links(ip: IPRoute = Depends(get_iproute)) -> List[InterfaceInfo]: + """List links in a normalized structure for UI consumers.""" + result: list[InterfaceInfo] = [] links = ip.get_links() addresses = ip.get_addr() - # group addresses by interface index - addr_map = {} - for a in addresses: - ifindex = a.get("index") - addr_map.setdefault(ifindex, []).append(a) + addr_map: dict[int, list] = {} + for addr in addresses: + ifindex = addr.get("index") + addr_map.setdefault(ifindex, []).append(addr) for link in links: - attrs = dict(link.get("attrs", [])) # convert list of tuples to dict + attrs = dict(link.get("attrs", [])) ifindex = link["index"] addrs = addr_map.get(ifindex, []) @@ -286,116 +254,95 @@ def get_raw_links(ip: IPRoute = Depends(get_iproute)): state=attrs.get("IFLA_OPERSTATE", "unknown"), mac=attrs.get("IFLA_ADDRESS"), mtu=attrs.get("IFLA_MTU", 0), - flags=[], # latest pyroute2 removed ifi_flags, leave empty + flags=[], addresses=parse_addresses(addrs), ) ) return result + @router.get("/bridges", response_model=List[BridgeInfo]) -def get_bridges(): - """ - Get all bridge interfaces on the system, including their member interfaces. - Returns detailed information: - - Bridge index, name, state, MTU - - STP state - - Member interfaces with index, name, state, and MTU - """ - bridges_list: List[BridgeInfo] = [] +def get_bridges() -> List[BridgeInfo]: + """List all bridges and their current member interfaces.""" + bridges_list: list[BridgeInfo] = [] - with NDB() as ndb: - for br in ndb.interfaces: - # Only bridges - if getattr(br, "kind", None) == "bridge": - members: List[BridgeInterfaceInfo] = [] - # Find member interfaces - for iface in ndb.interfaces: - if getattr(iface, "master", None) == br.index: - members.append( - BridgeInterfaceInfo( - ifindex=iface.index, - ifname=iface.ifname, - state=getattr(iface, "operstate", None), - mtu=getattr(iface, "mtu", None) - ) + with NDB() as ndb_ctx: + for bridge in ndb_ctx.interfaces: + if getattr(bridge, "kind", None) != "bridge": + continue + + members: list[BridgeInterfaceInfo] = [] + for iface in ndb_ctx.interfaces: + if getattr(iface, "master", None) == bridge.index: + members.append( + BridgeInterfaceInfo( + ifindex=iface.index, + ifname=iface.ifname, + state=getattr(iface, "operstate", 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 + @router.get("/full-state") -def full_state( - ip: IPRoute = Depends(get_iproute), -): - """ - Returns the full network state: - - Interfaces with IP addresses and flags - - Routes - - Bridges with member interfaces - """ +def full_state(ip: IPRoute = Depends(get_iproute)) -> dict: + """Return interfaces, routes, and bridges in one response.""" return { "interfaces": get_interfaces(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") 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, state="up") - # Interfaces hinzufügen + aktivieren for iface in req.interfaces: idx = iface_index(iface, ip) - - # interface hochfahren - ip.link("set", index=idx, state="down") # optional - sicherer + ip.link("set", index=idx, state="down") ip.link("set", index=idx, state="up") - - # interface in die bridge hängen ip.link("set", index=idx, master=br_idx) return { "status": "ok", "bridge": req.name, - "interfaces": req.interfaces + "interfaces": req.interfaces, } @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): - 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) - - # Bridge runterfahren ip.link("set", index=br_idx, state="down") - - # Bridge löschen ip.link("del", index=br_idx) return { "status": "ok", - "deleted": req.name - } \ No newline at end of file + "deleted": req.name, + } diff --git a/backend/src/api/packet_api.py b/backend/src/api/packet_api.py index ea06b50..cad84a4 100644 --- a/backend/src/api/packet_api.py +++ b/backend/src/api/packet_api.py @@ -1,96 +1,72 @@ -# src/routers/packets.py +"""Packet history and streaming endpoints.""" + import asyncio import base64 import json 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 src.Models.packets import PacketDBModel -from fastapi import APIRouter, Query, WebSocket, WebSocketDisconnect, HTTPException -from fastapi.responses import JSONResponse - import src.shared_objects as shared +from src.Models.packets import PacketDBModel logger = logging.getLogger("packets_router") router = APIRouter() def _serialize_row_for_json(row: Union[Dict[str, Any], PacketDBModel, BaseModel]) -> Dict[str, Any]: - """ - 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 + """Convert one packet row to a JSON-safe dictionary.""" 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: - # copy to avoid mutating caller's dict - d = dict(row) + raw_dict = dict(row) - # If raw_b64 already present, prefer it. If raw present and bytes, convert. - raw_val = d.get("raw") + raw_val = raw_dict.get("raw") if raw_val is not None and isinstance(raw_val, (bytes, bytearray)): - # convert to base64 string and remove raw try: - d["raw_b64"] = base64.b64encode(raw_val).decode("ascii") - d.pop("raw", None) + raw_dict["raw_b64"] = base64.b64encode(raw_val).decode("ascii") + raw_dict.pop("raw", None) except Exception: - # keep raw as str fallback try: - d["raw_b64"] = base64.b64encode(bytes(raw_val)).decode("ascii") - d.pop("raw", None) + raw_dict["raw_b64"] = base64.b64encode(bytes(raw_val)).decode("ascii") + raw_dict.pop("raw", None) except Exception: - logger.exception("Failed to base64-encode raw bytes for row id=%s", d.get("id")) - d["raw_b64"] = str(raw_val) - d.pop("raw", None) + logger.exception("Failed to base64-encode raw bytes for row id=%s", raw_dict.get("id")) + raw_dict["raw_b64"] = str(raw_val) + raw_dict.pop("raw", None) - # Ensure final dict is JSON-safe: try json.dumps on each value, fallback to str() - out: Dict[str, Any] = {} - for k, v in d.items(): - # 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)): + output: Dict[str, Any] = {} + for key, value in raw_dict.items(): + if key == "raw_b64" and isinstance(value, (bytes, bytearray)): try: - out["raw_b64"] = base64.b64encode(v).decode("ascii") - continue + output["raw_b64"] = base64.b64encode(value).decode("ascii") except Exception: - out["raw_b64"] = str(v) - continue + output["raw_b64"] = str(value) + continue - # JSON-serializable check try: - json.dumps({k: v}) - out[k] = v + json.dumps({key: value}) + output[key] = value except (TypeError, ValueError): - # convert non-serializable to string representation try: - out[k] = str(v) + output[key] = str(value) except Exception: - out[k] = "" - return out + output[key] = "" + + return output async def _serialize_rows(rows: List[Union[Dict[str, Any], PacketDBModel]]) -> List[Dict[str, Any]]: - """ - Serialize a list of DB rows or PacketDBModel instances into JSON-ready dicts. - Keeps the same order as input. - """ - return [_serialize_row_for_json(r) for r in rows] + """Convert packet rows to JSON-safe dictionaries, preserving order.""" + return [_serialize_row_for_json(row) for row in rows] + @router.get("/packets") -async def get_packets(limit: int = Query(100, ge=1, le=10000)): - """ - Return latest `limit` packets (newest first). The DB helper already converts - `raw` to `raw_b64` in fetch_latest, but we defensively re-serialize here. - """ +async def get_packets(limit: int = Query(100, ge=1, le=10000)) -> JSONResponse: + """Return the latest packets in reverse chronological order.""" db = shared.db if db is None: 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: rows = await db.fetch_latest(limit) - serial = await _serialize_rows(rows) - return JSONResponse(content={"count": len(serial), "packets": serial}) - except Exception: + serialized = await _serialize_rows(rows) + return JSONResponse(content={"count": len(serialized), "packets": serialized}) + except Exception as exc: 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") -async def websocket_packets(ws: WebSocket): - """ - 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. - """ +async def websocket_packets(ws: WebSocket) -> None: + """Stream live packets to a websocket client.""" await ws.accept() 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") return - # Parse subscribe_recent from query params (defensive) try: subscribe_recent_raw = ws.query_params.get("subscribe_recent", "0") subscribe_recent = int(subscribe_recent_raw) - if subscribe_recent < 0: - subscribe_recent = 0 + subscribe_recent = max(subscribe_recent, 0) except Exception: subscribe_recent = 0 - q: Optional[asyncio.Queue] = None + queue: Optional[asyncio.Queue] = None try: - # Optionally send recent history first if subscribe_recent > 0: recent = await db.fetch_latest(subscribe_recent) - recent_serial = await _serialize_rows(recent) - await ws.send_json({"type": "recent", "count": len(recent_serial), "packets": recent_serial}) + recent_serialized = await _serialize_rows(recent) + await ws.send_json({"type": "recent", "count": len(recent_serialized), "packets": recent_serialized}) - # Subscribe to broadcaster to receive live packets - q = await broadcaster.subscribe() - logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, q.maxsize) + queue = await broadcaster.subscribe() + logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, queue.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: - msg = await q.get() - # Normalize message to JSON-able dict - if isinstance(msg, dict): - payload = _serialize_row_for_json(msg) + message = await queue.get() + + if isinstance(message, dict): + payload: Any = _serialize_row_for_json(message) else: - # not a dict — try to json-serialize directly try: - json.dumps(msg) - payload = msg + json.dumps(message) + payload = message except Exception: - payload = {"data": str(msg)} + payload = {"data": str(message)} try: await ws.send_json(payload) 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 except WebSocketDisconnect: logger.info("WebSocket client disconnected: %s", ws.client) except Exception: logger.exception("Unexpected error in websocket_packets") finally: - # Clean up subscriber queue - if q is not None: + if queue is not None: try: - await broadcaster.unsubscribe(q) + await broadcaster.unsubscribe(queue) except Exception: logger.exception("Failed to unsubscribe websocket queue") + try: await ws.close() except Exception: 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") -async def clear_packets(reset_id: bool = Query(True)): - """ - Clear all packet logs from the database. - Uses TRUNCATE internally for high performance. - """ +async def clear_packets(reset_id: bool = Query(True)) -> JSONResponse: + """Remove all packet rows from the database.""" db = shared.db if db is None: raise HTTPException(status_code=503, detail="Database not available") success = await db.clear_all_packets(reset_identity=reset_id) - if not success: raise HTTPException(status_code=500, detail="Failed to clear packet table") logger.info("User initiated clear_packets (reset_id=%s)", reset_id) return JSONResponse( content={ - "status": "success", + "status": "success", "message": "All packets have been cleared", - "reset_id": reset_id + "reset_id": reset_id, } ) diff --git a/backend/src/api/sniffer_api.py b/backend/src/api/sniffer_api.py index 4d07104..0819e9c 100644 --- a/backend/src/api/sniffer_api.py +++ b/backend/src/api/sniffer_api.py @@ -1,7 +1,9 @@ -# src/routers/sniffer.py -from fastapi import APIRouter, HTTPException, Query, Body +"""HTTP API for starting, stopping, and inspecting sniffer sessions.""" + +from typing import Any, Dict, Optional + +from fastapi import APIRouter, Body, HTTPException, Query from pydantic import BaseModel, Field -from typing import Dict, Any, Optional from src.network_sniffer import ( get_sniffer_status, @@ -12,103 +14,113 @@ from src.network_sniffer import ( router = APIRouter() -# ------------------------------ -# Pydantic Models -# ------------------------------ class SnifferStartRequest(BaseModel): - """ - 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") + """Request payload for starting a sniffer session.""" - -class SnifferStartResponse(BaseModel): - """ - Response model returned when sniffer starts successfully. - """ - started: bool = Field(..., description="Whether the sniffer was started successfully") - session_id: str = Field(..., description="Session identifier for this sniffer instance") - target: str = Field(..., description="Target that was started (bridge or interface)") - target_type: str = Field(..., description="Either 'bridge' or 'interface'") - - -class SnifferStopRequest(BaseModel): - """ - Optional body for stop — prefer session_id if you want to stop a specific session. - If omitted, stopping behavior will be determined by query params (bridge/interface) or global stop. - """ - session_id: Optional[str] = Field(None, description="Session id to stop") - - -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)") - target: Optional[str] = Field(None, description="Target stopped; null if global stop") - target_type: Optional[str] = Field(None, description="'bridge' or 'interface' or None") - - -class InterfaceSnifferStatus(BaseModel): - running: bool = Field(..., description="Whether the sniffer thread/socket is active") - exists: bool = Field(..., description="Whether the interface exists in /sys/class/net") - up: bool = Field(..., description="Whether the interface is operationally UP") - session_id: Optional[str] = Field(None, description="Session id owning this interface") - session_label: Optional[str] = Field(None, description="Human label for the session") - - -class SnifferStatusResponse(BaseModel): - interfaces: Dict[str, InterfaceSnifferStatus] = Field( - ..., description="Map of interface names to their sniffer status" + bridge: Optional[str] = Field( + None, + example="br0", + description="Bridge name to sniff.", + ) + interface: Optional[str] = Field( + None, + example="eth0", + description="Interface name to sniff.", + ) + + +class SnifferStartResponse(BaseModel): + """Response payload for a successful sniffer start.""" + + started: bool = Field(..., description="True when a session was started.") + 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 SnifferStopRequest(BaseModel): + """Optional stop payload for targeting a specific session.""" + + session_id: Optional[str] = Field(None, description="Session ID to stop.") + + +class SnifferStopResponse(BaseModel): + """Response payload for stop operations.""" + + stopped: bool = Field(..., description="True when stop completed.") + 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 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) -def sniffer_start(req: SnifferStartRequest): - """ - 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. - """ +def sniffer_start(req: SnifferStartRequest) -> SnifferStartResponse: + """Start one sniffer session for exactly one target.""" if bool(req.bridge) == bool(req.interface): raise HTTPException(status_code=400, detail="Exactly one of 'bridge' or 'interface' must be provided") try: if req.interface: 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") - else: - 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") + return SnifferStartResponse( + started=True, + session_id=session_id, + 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: - 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) def sniffer_stop( - q_bridge: Optional[str] = Query(None, alias="bridge", description="If provided, stop sniffer sockets for this bridge"), - q_interface: Optional[str] = Query(None, alias="interface", description="If provided, stop sniffer socket for this interface"), + q_bridge: Optional[str] = Query( + 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(...), -): - """ - 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). - """ +) -> SnifferStopResponse: + """Stop by session ID, target query, or globally when no selector is given.""" if body and body.session_id: try: stop_afpacket_sniffer(session_id=body.session_id) return SnifferStopResponse(stopped=True, session_id=body.session_id, target=None, target_type=None) 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: 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: stop_afpacket_sniffer(target=q_interface, target_is_interface=True) return SnifferStopResponse(stopped=True, session_id=None, target=q_interface, target_type="interface") + if q_bridge: stop_afpacket_sniffer(target=q_bridge, target_is_interface=False) return SnifferStopResponse(stopped=True, session_id=None, target=q_bridge, target_type="bridge") - # global stop stop_afpacket_sniffer() return SnifferStopResponse(stopped=True, session_id=None, target=None, target_type=None) 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) -def sniffer_status(): - """ - Return the sniffer status information. - """ +def sniffer_status() -> SnifferStatusResponse: + """Return current sniffer status per interface.""" try: - raw = get_sniffer_status() - # Convert raw dict → typed model - typed = { - k: InterfaceSnifferStatus(**v) - for k, v in raw.items() - } + raw: Dict[str, Dict[str, Any]] = get_sniffer_status() + typed = {key: InterfaceSnifferStatus(**value) for key, value in raw.items()} return SnifferStatusResponse(interfaces=typed) except Exception as exc: - raise HTTPException(status_code=500, detail=f"Failed to query sniffer status: {exc}") \ No newline at end of file + raise HTTPException(status_code=500, detail=f"Failed to query sniffer status: {exc}") from exc diff --git a/backend/src/main.py b/backend/src/main.py index 1f99e88..79b98ec 100644 --- a/backend/src/main.py +++ b/backend/src/main.py @@ -1,30 +1,24 @@ -# src/main.py +"""FastAPI application entrypoint and runtime wiring.""" + import asyncio import logging import os + from fastapi import FastAPI 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.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 -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" - logging.basicConfig(level=logging.DEBUG) - -# ---- Globals ----------------------------------------------- -# Create DatabasePool instance (pool created on startup) shared_objects.db = DatabasePool(DB_DSN) app = FastAPI( @@ -45,62 +39,44 @@ app.add_middleware( allow_headers=["*"], ) -# --------------------- -# Startup / Shutdown -# --------------------- @app.on_event("startup") -async def on_startup(): - """ - Initialize DB pool and broadcaster on the FastAPI event loop and - publish them into shared_objects so other modules (sniffer, routers) - can access them. - """ +async def on_startup() -> None: + """Initialize shared runtime objects on the FastAPI event loop.""" loop = asyncio.get_running_loop() shared_objects.web_loop = loop - # Initialize DB pool bound to this loop try: await shared_objects.db.init_pool() except Exception: logging.exception("Failed to initialize DB pool") raise - # Create broadcaster and attach to DB so DB.insert_packet can publish updates try: shared_objects.broadcaster = PacketBroadcaster(loop) shared_objects.db.broadcaster = shared_objects.broadcaster except Exception: 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: - # import sniffer here to avoid circular imports at module import time from src import network_sniffer as sniffer - # sniffer provides drain_buffer_to_shared_db() try: sniffer.drain_buffer_to_shared_db() except Exception: logging.exception("Failed to drain sniffer buffer") 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") -async def shutdown_event(): - """ - Shutdown actions: stop network API and close DB pool if present. - """ - # try to shut down network API components +async def shutdown_event() -> None: + """Stop network resources and release shared runtime objects.""" try: network_api.shutdown_network_api() except Exception: logging.exception("Error shutting down network API") - # close DB pool if available in shared_objects try: web_db = getattr(shared_objects, "db", None) if web_db is not None: @@ -108,37 +84,26 @@ async def shutdown_event(): except Exception: logging.exception("Failed to close DB pool during shutdown") - # clear shared runtime objects (optional cleanup) - try: - shared_objects.db = None - shared_objects.broadcaster = None - shared_objects.web_loop = None - except Exception: - pass + shared_objects.db = None + shared_objects.broadcaster = None + shared_objects.web_loop = None -# --------------------- -# Basic Endpoints -# --------------------- - @app.get("/hello") -def hello(): +def hello() -> dict[str, str]: + """Simple health-check endpoint.""" return {"message": "Hello from FastAPI 🎉"} @app.get("/versions") -def versions(): +def versions() -> dict[str, str]: + """Return runtime Python version.""" message = os.popen("python --version").read().strip() return {"message": message} -# --------------------- -# Routers -# --------------------- app.include_router(network_api.router, prefix="/network", tags=["network"]) app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"]) app.include_router(packet_api.router, prefix="/packets", tags=["packets"]) -#app.include_router(nftables_api.router, prefix="/nftables", tags=["nftables"]) -#app.include_router(nft_api.router, prefix="/nft", tags=["nft"]) app.include_router(nft_manager.router, tags=["firewall"]) -app.include_router(packet_scripting_api.router, prefix="/scripts", tags=["scripts"]) \ No newline at end of file +app.include_router(packet_scripting_api.router, prefix="/scripts", tags=["scripts"]) diff --git a/backend/src/shared_objects.py b/backend/src/shared_objects.py index 552d362..2b3bcf4 100644 --- a/backend/src/shared_objects.py +++ b/backend/src/shared_objects.py @@ -1,8 +1,8 @@ -from typing import Optional -import asyncio +"""Shared runtime objects initialized during FastAPI startup.""" -# These are filled at FastAPI startup -# DB instance -db = None +import asyncio +from typing import Any, Optional + +db: Any = None web_loop: Optional[asyncio.AbstractEventLoop] = None -broadcaster = None \ No newline at end of file +broadcaster: Any = None diff --git a/backend/src/utilities/database.py b/backend/src/utilities/database.py index 6c872a4..5e1c298 100644 --- a/backend/src/utilities/database.py +++ b/backend/src/utilities/database.py @@ -1,28 +1,21 @@ -# src/utilities/database.py -import logging -import base64 -import asyncio -from typing import Dict, List, Optional, Any -from pydantic import ValidationError +"""Database helper for packet persistence and retrieval.""" +import asyncio +import base64 +import logging +from typing import Any, Dict, List, Optional import asyncpg from asyncpg.pool import Pool +from pydantic import ValidationError from src.Models.packets import PacketDBModel -# ---- Logging ---------------------------------------------------------- logger = logging.getLogger("af_packet_sniffer") class DatabasePool: - """ - 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 - """ + """Asyncpg connection pool wrapper used by the packet APIs.""" def __init__(self, dsn: str, min_size: int = 1, max_size: int = 5): self._dsn = dsn @@ -30,26 +23,26 @@ class DatabasePool: self._min_size = min_size self._max_size = max_size self.broadcaster = None - # created on first init_pool() (must be created on an event loop) self._init_lock: Optional[asyncio.Lock] = 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: return - # Ensure a lock exists that is bound to the running event loop if self._init_lock is None: self._init_lock = asyncio.Lock() async with self._init_lock: - # Double-check after acquiring lock if self._pool is not None: return + logger.info("Initializing DB pool (dsn=%s)", self._dsn) try: 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") except Exception: @@ -57,9 +50,10 @@ class DatabasePool: raise async def close_pool(self) -> None: - """Close the pool if it exists.""" + """Close the pool if present.""" if self._pool is None: return + try: await self._pool.close() logger.info("DB pool closed") @@ -69,13 +63,7 @@ class DatabasePool: self._pool = None async def insert_packet(self, pkt_info: Dict[str, Any]) -> None: - """ - 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). - """ + """Insert one packet record and publish it to subscribers.""" if self._pool is None: await self.init_pool() @@ -115,13 +103,11 @@ class DatabasePool: except Exception: logger.exception("DB insert failed") return - # Update the dictionary with the DB-generated values + if new_row: pkt_info["id"] = new_row["id"] - # Convert timestamp to ISO string for JSON serialization in WebSockets pkt_info["timestamp"] = new_row["timestamp"].isoformat() - # notify broadcaster (non-blocking). broadcaster is expected to be thread-safe. if self.broadcaster: try: self.broadcaster.sync_publish(pkt_info) @@ -129,11 +115,7 @@ class DatabasePool: logger.exception("Failed to publish pkt_info to broadcaster") async def fetch_latest(self, limit: int) -> List[PacketDBModel]: - """ - Fetch the latest `limit` packets (newest first). - - Returns a list of PacketDBModel. Converts raw bytes -> raw_b64 for JSON-safe output. - """ + """Fetch newest packet rows as validated `PacketDBModel` instances.""" if self._pool is None: await self.init_pool() @@ -148,45 +130,34 @@ class DatabasePool: limit, ) - out: List[PacketDBModel] = [] + result: List[PacketDBModel] = [] + for row in rows: + data = dict(row) - for r in rows: - d = dict(r) - - # convert byte raw -> base64 string (and remove raw) - raw_val = d.get("raw") + raw_val = data.get("raw") if isinstance(raw_val, (bytes, bytearray)): - d["raw_b64"] = base64.b64encode(raw_val).decode("ascii") - d.pop("raw", None) + data["raw_b64"] = base64.b64encode(raw_val).decode("ascii") + data.pop("raw", None) - - # Validate/construct Pydantic model try: - packet_model = PacketDBModel(**d) - except ValidationError as ve: - # Log and skip invalid rows (or handle otherwise) + packet_model = PacketDBModel(**data) + except ValidationError as exc: logger.warning( "Skipping DB row that failed PacketDBModel validation (id=%s): %s", - d.get("id"), - ve, + data.get("id"), + exc, ) continue - out.append(packet_model) + result.append(packet_model) + + return result - return out - async def clear_all_packets(self, reset_identity: bool = True) -> bool: - """ - 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. - """ + """Truncate the packet table and optionally reset identity counters.""" if self._pool is None: await self.init_pool() - # TRUNCATE is faster than DELETE and resets the identity counter restart_clause = "RESTART IDENTITY" if reset_identity else "" query = f"TRUNCATE TABLE packets {restart_clause};" diff --git a/backend/src/utilities/packet_broadcaster.py b/backend/src/utilities/packet_broadcaster.py index daa3d7a..d4d7db0 100644 --- a/backend/src/utilities/packet_broadcaster.py +++ b/backend/src/utilities/packet_broadcaster.py @@ -1,140 +1,100 @@ -# src/utilities/packet_broadcaster.py +"""In-process packet broadcaster for websocket subscribers.""" + import asyncio import logging -from typing import Dict, Any, List, Optional +from typing import Any, Dict, List, Optional logger = logging.getLogger("packet_broadcaster") class PacketBroadcaster: - """ - 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. - """ + """Manage subscriber queues and publish packet events.""" def __init__(self, loop: asyncio.AbstractEventLoop, queue_maxsize: int = 1024): self._loop = loop self._queue_maxsize = queue_maxsize - - # create lock and subscribers on the target loop to avoid cross-loop asyncio primitives 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._closed = False + try: - # ensure lock is created on the given loop - def _make_lock(): + def _make_lock() -> None: self._lock = asyncio.Lock() loop.call_soon_threadsafe(_make_lock) except Exception: - # fallback — create in current loop if call_soon_threadsafe fails self._lock = asyncio.Lock() - self._closed = False - async def subscribe(self) -> asyncio.Queue: - """ - Create a subscriber queue and add it to the list. - Caller is expected to await on the returned queue to receive messages. - """ + """Create and register a queue for one subscriber.""" if self._closed: raise RuntimeError("PacketBroadcaster is closed") - q: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize) - # wait until lock exists + queue: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize) while self._lock is None: - await asyncio.sleep(0) # yield to event loop briefly + await asyncio.sleep(0) async with self._lock: - self._subscribers.append(q) - return q + self._subscribers.append(queue) - async def unsubscribe(self, q: asyncio.Queue) -> None: - """ - Remove a subscriber queue if present. - """ + return queue + + async def unsubscribe(self, queue: asyncio.Queue) -> None: + """Unregister a subscriber queue if it exists.""" if self._lock is None: return + async with self._lock: try: - self._subscribers.remove(q) + self._subscribers.remove(queue) except ValueError: pass async def publish(self, msg: Dict[str, Any]) -> None: - """ - Publish msg to all subscribers (must be called on the broadcaster's loop). - 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 + """Publish one message to all current subscribers.""" + if self._closed or self._lock is None: return async with self._lock: - subs = list(self._subscribers) + subscribers = list(self._subscribers) - for q in subs: + for queue in subscribers: try: - q.put_nowait(msg) + queue.put_nowait(msg) except asyncio.QueueFull: - # drop message for this subscriber continue except Exception as exc: - logger.exception("Unexpected error when publishing to subscriber: %s", exc) - # attempt to remove broken subscriber + logger.exception("Unexpected subscriber publish error: %s", exc) try: async with self._lock: - if q in self._subscribers: - self._subscribers.remove(q) + if queue in self._subscribers: + self._subscribers.remove(queue) except Exception: pass def sync_publish(self, msg: Dict[str, Any]) -> None: - """ - 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. - """ + """Thread-safe wrapper that schedules `publish` on the broadcaster loop.""" if self._closed: return try: - # schedule the coroutine to run on the broadcaster loop self._loop.call_soon_threadsafe(asyncio.create_task, self.publish(msg)) except Exception as exc: - # swallow errors but log for debugging logger.exception("sync_publish failed to schedule publish: %s", exc) async def close(self) -> None: - """ - Close the broadcaster: mark closed, clear subscribers, and drain queues. - """ + """Close the broadcaster and clear queued messages.""" self._closed = True if self._lock is None: return + async with self._lock: - subs = list(self._subscribers) + subscribers = list(self._subscribers) self._subscribers.clear() - for q in subs: + for queue in subscribers: try: - # optionally notify subscribers of closure by putting None (client must handle) - # q.put_nowait(None) - while not q.empty(): - try: - q.get_nowait() - except Exception: - break + while not queue.empty(): + queue.get_nowait() except Exception: pass diff --git a/frontend/src/api/apiClient.ts b/frontend/src/api/apiClient.ts index 3a53d06..8af0f9d 100644 --- a/frontend/src/api/apiClient.ts +++ b/frontend/src/api/apiClient.ts @@ -1,5 +1,5 @@ -// src/apiClient.ts import axios from 'axios'; + import { CreateRuleRequest, ExecResult, RulesetModel } from '../types/firewall'; import { BridgeCreateRequest, @@ -36,22 +36,14 @@ export const api = axios.create({ timeout: 20000, }); -// Normalize FastAPI errors here api.interceptors.response.use( (response) => response, (error) => { - // FastAPI HTTPException format 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)); }, ); -/* ------------------------- - Basic endpoints - ------------------------- */ - export const fetchHello = async (): Promise => { const res = await api.get('/hello'); return res.data; @@ -62,10 +54,6 @@ export const fetchVersions = async (): Promise => { return res.data; }; -/* ------------------------- - Network queries (existing) - ------------------------- */ - export const fetchInterfaces = async (): Promise => { const res = await api.get('/network/interfaces'); return res.data; @@ -101,150 +89,84 @@ export const removeBridge = async (req: BridgeRemoveRequest) => { 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 => { const res = await api.post('/sniffer/start', payload); 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 => { const res = await api.post('/sniffer/stop', body ?? {}); 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 => { const res = await api.post(`/sniffer/stop?interface=${encodeURIComponent(iface)}`, {}); 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 => { const res = await api.post(`/sniffer/stop?bridge=${encodeURIComponent(bridge)}`, {}); return res.data; }; -/** - * Fetch the sniffer status (per-interface). - */ export const fetchSnifferStatus = async (): Promise => { const res = await api.get('/sniffer/status'); return res.data; }; -/* ------------------------- - Packets - ------------------------- */ export const fetchPackets = async (limit = 100): Promise => { - // limit default mirrors OpenAPI default const res = await api.get('/packets/packets', { params: { limit } }); return res.data; }; export const clearPackets = async (): Promise => { - // limit default mirrors OpenAPI default const res = await api.delete('/packets/packets'); 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 }> => { const res = await api.get<{ ruleset: RulesetModel }>('/firewall/rules'); 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 => { const res = await api.delete(`/firewall/rules/${encodeURIComponent(String(handle))}`, { params: { family, table, chain }, }); - // axios resolves non-2xx as reject; server uses 204 No Content so nothing to return return res.data; }; -/** - * createRuleJson - POST /firewall/rules - * Body: CreateRuleRequest (must include expr) - */ export const createRuleJson = async (req: CreateRuleRequest): Promise => { const res = await api.post('/firewall/rules', req); return res.data; }; export const execFirewallRaw = async (cmd: string): Promise => { - const res = await api.post('/firewall/raw', { cmd: cmd }); + const res = await api.post('/firewall/raw', { cmd }); return res.data; }; -/* ------------------------- - Scripts - ------------------------- */ - -/** - * Fetch the combined scripts + status endpoint. - * Returns a list of ScriptWithStatus entries. - */ export const fetchScriptsAll = async (): Promise => { const res = await api.get('/scripts/scripts'); 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 => { 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 => { const all = await fetchScriptsAll(); - const found = all.find((s) => s.name === name); + const found = all.find((script) => script.name === name); + if (!found) { - // If script not found we still return an empty mapping structure return { name, 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: { name: string; script: File | Blob; @@ -253,7 +175,10 @@ export const uploadScript = async (opts: { const fd = new FormData(); fd.append('name', opts.name); 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('/scripts/scripts', fd, { headers: { 'Content-Type': 'multipart/form-data' }, @@ -261,71 +186,49 @@ export const uploadScript = async (opts: { 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 => { const res = await api.get(`/scripts/scripts/${encodeURIComponent(name)}`, { responseType: '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 => { const res = await api.get(`/scripts/scripts/${encodeURIComponent(name)}/requirements`, { responseType: '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 ( name: string, requirements: File | Blob, ): Promise => { const fd = new FormData(); fd.append('requirements', requirements); + const res = await api.put(`/scripts/scripts/${encodeURIComponent(name)}/requirements`, fd, { headers: { 'Content-Type': 'multipart/form-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 => { const res = await api.delete(`/scripts/scripts/${encodeURIComponent(name)}/requirements`); 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 => { const params: Record = {}; - if (typeof qnum !== 'undefined' && qnum !== null) params.qnum = qnum; + if (typeof qnum !== 'undefined' && qnum !== null) { + params.qnum = qnum; + } + const res = await api.delete(`/scripts/scripts/${encodeURIComponent(name)}`, { params }); 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 => { const res = await api.post(`/scripts/scripts/${encodeURIComponent(name)}/enable`, req); return res.data; }; -/** - * Disable a script unit (POST with qnum as query param). Returns OperationResult. - */ export const disableScript = async (name: string, qnum: number): Promise => { const res = await api.post(`/scripts/scripts/${encodeURIComponent(name)}/disable`, null, { params: { qnum }, diff --git a/frontend/src/appRouter.tsx b/frontend/src/appRouter.tsx index ac5bd1c..46ef434 100644 --- a/frontend/src/appRouter.tsx +++ b/frontend/src/appRouter.tsx @@ -1,6 +1,6 @@ -// src/AppRouter.tsx import { Navigate, Route, Routes } from 'react-router-dom'; -import App from './App'; // your layout component (has ) + +import App from './App'; import { Firewall } from './pages/Firewall'; import Home from './pages/Home'; import Network from './pages/Network'; @@ -8,34 +8,27 @@ import Scripting from './pages/Scripting'; import Sniffing from './pages/Sniffing'; import { PATHS } from './routes'; +function NotFound() { + return ( +
+

404 - Not Found

+

The requested page does not exist.

+
+ ); +} + export default function AppRouter() { return ( - {/* App is the top-level layout; Outlet renders the active child route */} }> - {/* When the user hits '/', redirect to '/home' */} } /> - - {/* Child routes - these render inside App's */} } /> } /> } /> } /> } /> - - {/* Fallback (renders inside layout too) */} } /> ); } - -/** simple 404 rendered inside the layout */ -function NotFound() { - return ( -
-

404 – Not Found

-

The requested page does not exist.

-
- ); -} diff --git a/frontend/src/components/FireWallAddChainModal.tsx b/frontend/src/components/FireWallAddChainModal.tsx index 9005f35..f4341be 100644 --- a/frontend/src/components/FireWallAddChainModal.tsx +++ b/frontend/src/components/FireWallAddChainModal.tsx @@ -1,4 +1,3 @@ -// src/components/AddChainModal.tsx import { CopyOutlined } from '@ant-design/icons'; import { Alert, @@ -67,12 +66,10 @@ export default function FirewallAddChainModal({ const isPrefilled = Boolean(table?.family && table?.name); - // init form values when modal opens or table prop changes useEffect(() => { form.setFieldsValue({ family: table?.family ?? 'bridge', tableName: table?.name ?? 'filter', - // Do not override chainName if the user has typed it previously type: 'filter', hook: 'forward', 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[] { const vals = values ?? form.getFieldsValue(); @@ -128,7 +124,6 @@ export default function FirewallAddChainModal({ : 0; const policy = vals.policy ?? ''; - // chain name: prefer explicit chainName else fallback to hook (less ideal) else 'mychain' const chain = hook; const policyPart = policy ? ` policy ${policy} ;` : ''; @@ -136,7 +131,6 @@ export default function FirewallAddChainModal({ return [cmd]; } - // Execute commands sequentially async function executeCommands(cmds: string[]) { setRunning(true); setResults([]); @@ -156,24 +150,19 @@ export default function FirewallAddChainModal({ setResults(acc); setRunning(false); - // refresh ruleset shown in modal try { await refreshRuleset(); } catch { - /* ignored, refreshRuleset already notified on failure */ } const hadError = acc.some((r) => r.err); if (!hadError) { - // success notification inside the modal (not global refresh notification) notification.success({ message: 'Chain created', description: 'Chain created and ruleset refreshed locally in the modal.', duration: 4, }); - // Inform parent that a resource was created. - // Parent can decide whether to refresh and whether to show a notification. onClose?.(true); if (onSuccess) onSuccess(); } else { @@ -219,13 +208,11 @@ export default function FirewallAddChainModal({ type="primary" onClick={async () => { try { - // validate name fields: chainName required; if not prefilled, tableName & family required const requiredFields = ['chainName']; if (!isPrefilled) requiredFields.push('tableName', 'family'); await form.validateFields(requiredFields as any); setStep(1); } catch { - // AntD will show validation messages; no extra handling required } }} > @@ -264,7 +251,7 @@ export default function FirewallAddChainModal({ footer={renderFooter()} destroyOnClose > - {/* Step 0: Form */} + {step === 0 && (
{rulesetEmpty && ( @@ -284,7 +271,7 @@ export default function FirewallAddChainModal({ }} > - {/* family & table: show inputs only when not provided via props */} + {isPrefilled ? ( @@ -363,7 +350,7 @@ export default function FirewallAddChainModal({
)} - {/* Step 1: Preview & Results */} + {step === 1 && ( <> diff --git a/frontend/src/components/FireWallAddTableModal.tsx b/frontend/src/components/FireWallAddTableModal.tsx index 5eceb74..eb8991a 100644 --- a/frontend/src/components/FireWallAddTableModal.tsx +++ b/frontend/src/components/FireWallAddTableModal.tsx @@ -1,4 +1,3 @@ -// src/components/FirewallManager.tsx import { Button, Card, Col, Form, Input, Modal, Row, Select, Space, Spin, Typography, notification } from 'antd'; import { ReactElement, useEffect, useState } from 'react'; import { execFirewallRaw, fetchRuleset } from '../api/apiClient'; @@ -42,7 +41,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview = const [results, setResults] = useState([]); const [selectedFamily, setSelectedFamily] = useState('bridge'); - // initialize and refresh when modal opens useEffect(() => { form.setFieldsValue({ family: 'bridge', @@ -70,7 +68,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview = setRulesetEmpty(false); } } catch (err: any) { - // use notification instead of message notification.warning({ message: 'Could not load ruleset', 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[] { const family = (values.family ?? 'bridge').trim(); const table = (values.tableName ?? 'filter').trim(); return [`add table ${family} ${table}`]; } - // Execute commands sequentially async function executeCommands(cmds: string[]) { setRunning(true); setResults([]); @@ -109,16 +104,13 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview = setResults(acc); setRunning(false); - // refresh ruleset after running try { await refreshRuleset(); } catch { - // ignore — refreshRuleset handles notifications on error } const hadError = acc.some((r) => r.err); if (!hadError) { - // success -> notify briefly and inform parent that a resource was created notification.success({ message: 'Table created', 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 } else { - // error -> open results view and show notification notification.error({ message: 'Some commands returned errors', description: 'See execution results below for details.', @@ -153,7 +144,6 @@ export default function FirewallAddTableModal({ open, onClose, startOnPreview = - + - + r.ifindex} + rowKey={(row: InterfaceInfo) => row.ifindex} dataSource={networkState?.interfaces ?? []} columns={interfaceColumns} pagination={{ pageSize: 8 }} @@ -198,14 +187,13 @@ export default function Network() { } type="primary" onClick={() => { - // ensure up-to-date interface list when opening modal getFullState(true); setBridgeModalVisible(true); }} @@ -213,7 +201,7 @@ export default function Network() { } >
String(r.ifindex)} + rowKey={(row: BridgeInfo) => String(row.ifindex)} dataSource={networkState?.bridges ?? []} columns={bridgeColumns} pagination={{ pageSize: 6 }} @@ -223,7 +211,6 @@ export default function Network() { - {/* Create Bridge Modal */} - {/* Select field populated from interfaces endpoint */}