use shared objects like loops db access etc
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
This commit is contained in:
144
backend/src/api/packet_api.py
Normal file
144
backend/src/api/packet_api.py
Normal file
@@ -0,0 +1,144 @@
|
||||
# src/routers/packets.py
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
from typing import Optional, Any, Dict, List
|
||||
|
||||
from fastapi import APIRouter, Query, WebSocket, WebSocketDisconnect, HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
import src.shared_objects as shared
|
||||
|
||||
logger = logging.getLogger("packets_router")
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _serialize_row_for_json(row: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert DB row / pkt_info to a JSON-serializable dict.
|
||||
- If 'raw' is bytes, produce 'raw_b64' and drop 'raw'.
|
||||
- Fallback to str() for unknown/unserializable values.
|
||||
"""
|
||||
out: Dict[str, Any] = {}
|
||||
for k, v in row.items():
|
||||
if k == "raw" and isinstance(v, (bytes, bytearray)):
|
||||
out["raw_b64"] = base64.b64encode(v).decode("ascii")
|
||||
continue
|
||||
# try to JSON serialize the value directly
|
||||
try:
|
||||
json.dumps({k: v})
|
||||
out[k] = v
|
||||
except (TypeError, ValueError):
|
||||
out[k] = str(v)
|
||||
return out
|
||||
|
||||
|
||||
async def _serialize_rows(rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
return [_serialize_row_for_json(r) for r 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.
|
||||
"""
|
||||
db = shared.db
|
||||
if db is None:
|
||||
logger.warning("GET /packets called but DB is not available")
|
||||
raise HTTPException(status_code=503, detail="Database not available")
|
||||
|
||||
try:
|
||||
rows = await db.fetch_latest(limit)
|
||||
serial = await _serialize_rows(rows)
|
||||
return JSONResponse(content={"count": len(serial), "packets": serial})
|
||||
except Exception:
|
||||
logger.exception("Failed to fetch latest packets from DB")
|
||||
raise HTTPException(status_code=500, detail="Failed to fetch packets")
|
||||
|
||||
|
||||
@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.
|
||||
"""
|
||||
await ws.accept()
|
||||
logger.debug("WebSocket connection accepted: %s", ws.client)
|
||||
|
||||
db = shared.db
|
||||
broadcaster = shared.broadcaster
|
||||
|
||||
if db is None:
|
||||
await ws.send_json({"error": "database not available"})
|
||||
await ws.close()
|
||||
logger.warning("WebSocket closed: DB not available")
|
||||
return
|
||||
|
||||
if broadcaster is None:
|
||||
await ws.send_json({"error": "broadcaster not available"})
|
||||
await ws.close()
|
||||
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
|
||||
except Exception:
|
||||
subscribe_recent = 0
|
||||
|
||||
q: 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})
|
||||
|
||||
# Subscribe to broadcaster to receive live packets
|
||||
q = await broadcaster.subscribe()
|
||||
logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, q.maxsize)
|
||||
|
||||
# Simple heartbeat: periodically ensure client is responsive (optional)
|
||||
# We'll implement by awaiting q.get() which blocks until a message is published.
|
||||
while True:
|
||||
msg = await q.get()
|
||||
# Normalize message to JSON-able dict
|
||||
if isinstance(msg, dict):
|
||||
payload = _serialize_row_for_json(msg)
|
||||
else:
|
||||
# not a dict — try to json-serialize directly
|
||||
try:
|
||||
json.dumps(msg)
|
||||
payload = msg
|
||||
except Exception:
|
||||
payload = {"data": str(msg)}
|
||||
|
||||
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)
|
||||
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:
|
||||
try:
|
||||
await broadcaster.unsubscribe(q)
|
||||
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)
|
||||
@@ -1,12 +1,23 @@
|
||||
# src/main.py
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
import os
|
||||
import src.Models.netplan as netplan
|
||||
import src.network_api as network_api
|
||||
import src.routes_sniffer as sniffer_router
|
||||
from asyncio import subprocess
|
||||
from fastapi import HTTPException
|
||||
|
||||
from src.api import packet_api
|
||||
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
|
||||
|
||||
# ---- Config -----------------------------------------------------------
|
||||
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
|
||||
|
||||
# ---- Globals -----------------------------------------------
|
||||
# Create DatabasePool instance (pool created on startup)
|
||||
shared_objects.db = DatabasePool(DB_DSN)
|
||||
|
||||
app = FastAPI(
|
||||
root_path="/api",
|
||||
@@ -26,15 +37,76 @@ 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.
|
||||
"""
|
||||
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 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")
|
||||
|
||||
|
||||
@app.on_event("shutdown")
|
||||
def shutdown_event():
|
||||
async def shutdown_event():
|
||||
"""
|
||||
Shutdown actions: stop network API and close DB pool if present.
|
||||
"""
|
||||
# try to shut down network API components
|
||||
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:
|
||||
await web_db.close_pool()
|
||||
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
|
||||
|
||||
|
||||
# ---------------------
|
||||
@@ -45,11 +117,13 @@ def shutdown_event():
|
||||
def hello():
|
||||
return {"message": "Hello from FastAPI 🎉"}
|
||||
|
||||
|
||||
@app.get("/versions")
|
||||
def versions():
|
||||
message = os.popen("python --version").read().strip()
|
||||
return {"message": message}
|
||||
|
||||
|
||||
@app.get("/nft/ruleset")
|
||||
def nft_ruleset():
|
||||
message = os.popen("sudo nft --json list ruleset").read().strip()
|
||||
@@ -61,4 +135,5 @@ def nft_ruleset():
|
||||
# ---------------------
|
||||
|
||||
app.include_router(network_api.router, prefix="/network", tags=["network"])
|
||||
app.include_router(sniffer_router.router, prefix="/sniffer", tags=["sniffer"])
|
||||
app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"])
|
||||
app.include_router(packet_api.router, prefix="/packets", tags=["packets"])
|
||||
@@ -1,3 +1,4 @@
|
||||
# src/sniffer.py
|
||||
import asyncio
|
||||
import logging
|
||||
import threading
|
||||
@@ -8,8 +9,9 @@ import errno
|
||||
import struct
|
||||
from typing import Dict, List, Optional, Any
|
||||
|
||||
import asyncpg
|
||||
from src.utilities.database import DB
|
||||
# NOTE: ensure this path points to your shared runtime module
|
||||
from backend.src import shared_objects
|
||||
|
||||
from scapy.all import (
|
||||
Ether,
|
||||
ARP,
|
||||
@@ -35,9 +37,6 @@ from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("af_packet_sniffer")
|
||||
|
||||
# ---- Config -----------------------------------------------------------
|
||||
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
|
||||
|
||||
# ---- Globals (kept minimal) -------------------------------------------
|
||||
af_sockets: Dict[str, socket.socket] = {}
|
||||
af_thread: Optional[threading.Thread] = None
|
||||
@@ -46,7 +45,11 @@ af_stop_event: Optional[threading.Event] = None
|
||||
fixed_bridge_ports: Dict[str, List[str]] = {}
|
||||
current_bridge: Optional[str] = None
|
||||
|
||||
# Background asyncio loop used to run DB tasks
|
||||
# small bounded buffer for packets produced before shared_objects is ready
|
||||
_PACKET_BUFFER: List[Dict[str, Any]] = []
|
||||
_BUFFER_CAPACITY = 20000
|
||||
|
||||
# Background asyncio loop used for internal tasks in this module (kept but not used for DB pool)
|
||||
async_loop = asyncio.new_event_loop()
|
||||
|
||||
|
||||
@@ -60,11 +63,37 @@ threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).star
|
||||
|
||||
|
||||
# -------------------------
|
||||
# Database init (pool)
|
||||
# Helpers for buffer draining
|
||||
# -------------------------
|
||||
db = DB(DB_DSN)
|
||||
# schedule pool init on background loop (fire-and-forget)
|
||||
asyncio.run_coroutine_threadsafe(db.init_pool(), async_loop)
|
||||
def drain_buffer_to_shared_db() -> None:
|
||||
"""
|
||||
Attempt to schedule buffered packets for insertion on shared_objects.web_loop.
|
||||
Call this from main.py after shared_objects.db and shared_objects.web_loop are initialized.
|
||||
"""
|
||||
try:
|
||||
web_loop = getattr(shared_objects, "web_loop", None)
|
||||
web_db = getattr(shared_objects, "db", None)
|
||||
if web_db is None or web_loop is None:
|
||||
return
|
||||
|
||||
# schedule draining on the web loop to avoid blocking this thread
|
||||
def _drain() -> None:
|
||||
while _PACKET_BUFFER:
|
||||
pkt = _PACKET_BUFFER.pop(0)
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(web_db.insert_packet(pkt), web_loop)
|
||||
except Exception:
|
||||
# re-buffer first element and stop to avoid busy loop
|
||||
_PACKET_BUFFER.insert(0, pkt)
|
||||
break
|
||||
|
||||
try:
|
||||
web_loop.call_soon_threadsafe(_drain)
|
||||
except Exception:
|
||||
# fallback: run directly (best-effort)
|
||||
_drain()
|
||||
except Exception:
|
||||
logger.exception("Failed to drain packet buffer")
|
||||
|
||||
|
||||
# -------------------------
|
||||
@@ -80,8 +109,6 @@ def _safe_get_attr(layer, attr: str):
|
||||
def parse_packet(pkt, bridge: str) -> None:
|
||||
"""
|
||||
Parse a scapy Packet object into a normalized dict and schedule DB insert.
|
||||
Keeps all information from the original implementation (fields, VLAN handling,
|
||||
IP/IPv6/TCP/UDP/ARP handling and enum mapping).
|
||||
"""
|
||||
pkt_iface = getattr(pkt, "sniffed_on", None)
|
||||
if not pkt_iface:
|
||||
@@ -216,9 +243,18 @@ def parse_packet(pkt, bridge: str) -> None:
|
||||
if Raw in pkt and not pkt_info["protocol_name"]:
|
||||
pkt_info["protocol_name"] = "RAW"
|
||||
|
||||
# Submit DB insert to background loop (non-blocking)
|
||||
# Submit DB insert to shared web loop if available, otherwise buffer
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(db.insert_packet(pkt_info), async_loop)
|
||||
web_loop = getattr(shared_objects, "web_loop", None)
|
||||
web_db = getattr(shared_objects, "db", None)
|
||||
if web_db is not None and web_loop is not None:
|
||||
asyncio.run_coroutine_threadsafe(web_db.insert_packet(pkt_info), web_loop)
|
||||
else:
|
||||
# buffer (bounded) until the web app initializes
|
||||
_PACKET_BUFFER.append(pkt_info)
|
||||
if len(_PACKET_BUFFER) > _BUFFER_CAPACITY:
|
||||
# drop oldest packet
|
||||
_PACKET_BUFFER.pop(0)
|
||||
except Exception:
|
||||
logger.exception("Failed to schedule DB insert")
|
||||
|
||||
@@ -419,9 +455,12 @@ def stop_afpacket_sniffer() -> None:
|
||||
fixed_bridge_ports.pop(current_bridge, None)
|
||||
current_bridge = None
|
||||
|
||||
# close db pool (schedule on async loop)
|
||||
# close db pool if shared.web_loop is available; otherwise leave to main
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(db.close_pool(), async_loop)
|
||||
web_loop = getattr(shared_objects, "web_loop", None)
|
||||
web_db = getattr(shared_objects, "db", None)
|
||||
if web_db is not None and web_loop is not None:
|
||||
asyncio.run_coroutine_threadsafe(web_db.close_pool(), web_loop)
|
||||
except Exception:
|
||||
logger.exception("Failed to schedule DB pool close")
|
||||
|
||||
|
||||
8
backend/src/shared_objects.py
Normal file
8
backend/src/shared_objects.py
Normal file
@@ -0,0 +1,8 @@
|
||||
from typing import Optional
|
||||
import asyncio
|
||||
|
||||
# These are filled at FastAPI startup
|
||||
# DB instance
|
||||
db = None
|
||||
web_loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
broadcaster = None
|
||||
@@ -1,47 +1,80 @@
|
||||
|
||||
# src/utilities/database.py
|
||||
import logging
|
||||
import base64
|
||||
import asyncio
|
||||
from typing import Dict, List, Optional, Any
|
||||
|
||||
import asyncpg
|
||||
from asyncpg.pool import Pool
|
||||
|
||||
# ---- Logging ----------------------------------------------------------
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("af_packet_sniffer")
|
||||
|
||||
class DB:
|
||||
|
||||
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
|
||||
"""
|
||||
|
||||
def __init__(self, dsn: str, min_size: int = 1, max_size: int = 5):
|
||||
"""
|
||||
Create a DB helper bound to the given DSN.
|
||||
Pool is created lazily on first use unless init_pool() is explicitly called.
|
||||
"""
|
||||
self._dsn = dsn
|
||||
self._pool: Optional[asyncpg.pool.Pool] = None
|
||||
self._pool: Optional[Pool] = None
|
||||
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:
|
||||
"""Create connection pool if not already created."""
|
||||
if self._pool is None:
|
||||
"""Initialize the asyncpg pool if not already initialized (idempotent)."""
|
||||
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:
|
||||
logger.exception("Failed to create DB pool")
|
||||
raise
|
||||
|
||||
async def close_pool(self) -> None:
|
||||
"""Gracefully close connection pool."""
|
||||
if self._pool:
|
||||
"""Close the pool if it exists."""
|
||||
if self._pool is None:
|
||||
return
|
||||
try:
|
||||
await self._pool.close()
|
||||
self._pool = None
|
||||
logger.info("DB pool closed")
|
||||
except Exception:
|
||||
logger.exception("Error closing DB pool")
|
||||
finally:
|
||||
self._pool = None
|
||||
|
||||
async def insert_packet(self, pkt_info: Dict[str, Any]) -> None:
|
||||
"""
|
||||
Insert packet metadata. Same fields + same SQL statement as before.
|
||||
Insert packet metadata into the `packets` table.
|
||||
|
||||
Preserves the same columns/values as before.
|
||||
After a successful insert, if a broadcaster is attached it will be
|
||||
notified via broadcaster.sync_publish(pkt_info).
|
||||
"""
|
||||
if self._pool is None:
|
||||
await self.init_pool() # lazy init
|
||||
await self.init_pool()
|
||||
|
||||
try:
|
||||
async with self._pool.acquire() as conn:
|
||||
@@ -77,3 +110,44 @@ class DB:
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("DB insert failed")
|
||||
return
|
||||
|
||||
# notify broadcaster (non-blocking). broadcaster is expected to be thread-safe.
|
||||
if self.broadcaster:
|
||||
try:
|
||||
self.broadcaster.sync_publish(pkt_info)
|
||||
except Exception:
|
||||
logger.exception("Failed to publish pkt_info to broadcaster")
|
||||
|
||||
async def fetch_latest(self, limit: int) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Fetch the latest `limit` packets (newest first).
|
||||
|
||||
Returns a list of dicts. If the `raw` column is binary it is converted
|
||||
to `raw_b64` (base64 string) and `raw` is removed.
|
||||
"""
|
||||
if self._pool is None:
|
||||
await self.init_pool()
|
||||
|
||||
async with self._pool.acquire() as conn:
|
||||
# Select explicit columns to ensure predictable dict keys
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, iface, src_mac, dst_mac, eth_type, vlan_id, src_ip, dst_ip,
|
||||
ip_proto, src_port, dst_port, length, raw, created_at
|
||||
FROM packets
|
||||
ORDER BY id DESC
|
||||
LIMIT $1
|
||||
""",
|
||||
limit,
|
||||
)
|
||||
|
||||
out: List[Dict[str, Any]] = []
|
||||
for r in rows:
|
||||
d = dict(r)
|
||||
raw_val = d.get("raw")
|
||||
if isinstance(raw_val, (bytes, bytearray)):
|
||||
d["raw_b64"] = base64.b64encode(raw_val).decode("ascii")
|
||||
d.pop("raw", None)
|
||||
out.append(d)
|
||||
return out
|
||||
|
||||
140
backend/src/utilities/packet_broadcaster.py
Normal file
140
backend/src/utilities/packet_broadcaster.py
Normal file
@@ -0,0 +1,140 @@
|
||||
# src/utilities/packet_broadcaster.py
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Dict, Any, 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.
|
||||
"""
|
||||
|
||||
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
|
||||
try:
|
||||
# ensure lock is created on the given loop
|
||||
def _make_lock():
|
||||
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.
|
||||
"""
|
||||
if self._closed:
|
||||
raise RuntimeError("PacketBroadcaster is closed")
|
||||
|
||||
q: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize)
|
||||
# wait until lock exists
|
||||
while self._lock is None:
|
||||
await asyncio.sleep(0) # yield to event loop briefly
|
||||
|
||||
async with self._lock:
|
||||
self._subscribers.append(q)
|
||||
return q
|
||||
|
||||
async def unsubscribe(self, q: asyncio.Queue) -> None:
|
||||
"""
|
||||
Remove a subscriber queue if present.
|
||||
"""
|
||||
if self._lock is None:
|
||||
return
|
||||
async with self._lock:
|
||||
try:
|
||||
self._subscribers.remove(q)
|
||||
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
|
||||
return
|
||||
|
||||
async with self._lock:
|
||||
subs = list(self._subscribers)
|
||||
|
||||
for q in subs:
|
||||
try:
|
||||
q.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
|
||||
try:
|
||||
async with self._lock:
|
||||
if q in self._subscribers:
|
||||
self._subscribers.remove(q)
|
||||
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.
|
||||
"""
|
||||
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.
|
||||
"""
|
||||
self._closed = True
|
||||
if self._lock is None:
|
||||
return
|
||||
async with self._lock:
|
||||
subs = list(self._subscribers)
|
||||
self._subscribers.clear()
|
||||
|
||||
for q in subs:
|
||||
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
|
||||
except Exception:
|
||||
pass
|
||||
Reference in New Issue
Block a user