file structure and comments unified
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s
This commit is contained in:
@@ -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
|
||||
}
|
||||
}
|
||||
}'''
|
||||
|
||||
@@ -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...",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
"deleted": req.name,
|
||||
}
|
||||
|
||||
@@ -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] = "<unserializable>"
|
||||
return out
|
||||
output[key] = "<unserializable>"
|
||||
|
||||
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,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -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}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to query sniffer status: {exc}") from exc
|
||||
|
||||
@@ -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"])
|
||||
app.include_router(packet_scripting_api.router, prefix="/scripts", tags=["scripts"])
|
||||
|
||||
@@ -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
|
||||
broadcaster: Any = None
|
||||
|
||||
@@ -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};"
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user