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

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

View File

@@ -1,13 +1,20 @@
"""Netplan schema models used by bridge/network configuration APIs."""
from typing import Dict, List, Optional
from pydantic import BaseModel, Field
from 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
}
}
}'''

View File

@@ -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...",
}
}
}

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

View File

@@ -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"])

View File

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

View File

@@ -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};"

View File

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