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

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

View File

@@ -1,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,
}

View File

@@ -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,
}
)

View File

@@ -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