diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index 1f9c2b8..1b93b60 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -6,9 +6,13 @@ import threading DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" -sniffer_tasks = {} sniffer_threads = {} +event_loop = None # <-- we store FastAPI's main loop here + +# ----------------------------- +# DB INSERT +# ----------------------------- async def db_insert_packet(pkt_info: dict): conn = await asyncpg.connect(DB_DSN) try: @@ -40,13 +44,15 @@ async def db_insert_packet(pkt_info: dict): await conn.close() +# ----------------------------- +# PACKET HANDLER (Thread) +# ----------------------------- def handle_packet(pkt, iface): - """ - Scapy callback → run in thread. - """ + global event_loop + pkt_info = { "iface": iface, - "direction": "unknown", # You can set eth0=ingress, eth1=egress + "direction": "unknown", "src_mac": pkt[Ether].src if Ether in pkt else None, "dst_mac": pkt[Ether].dst if Ether in pkt else None, "eth_type": pkt[Ether].type if Ether in pkt else None, @@ -54,26 +60,39 @@ def handle_packet(pkt, iface): "dst_ip": pkt[IP].dst if IP in pkt else None, "protocol": pkt[IP].proto if IP in pkt else None, "length": len(pkt), - "ebpf_verdict": None, # will integrate later + "ebpf_verdict": None, "ebpf_chain": None, - "raw": bytes(pkt) + "raw": bytes(pkt), } - # push to asyncio loop - loop = asyncio.get_event_loop() - loop.create_task(db_insert_packet(pkt_info)) + # Schedule coroutine in main event loop from thread + asyncio.run_coroutine_threadsafe( + db_insert_packet(pkt_info), + event_loop + ) +# ----------------------------- +# SNIFFER THREAD +# ----------------------------- def start_sniffer_thread(iface: str): def sniff_blocking(): - sniff(prn=lambda x: handle_packet(x, iface), iface=iface, store=False) + sniff(prn=lambda x: handle_packet(x, iface), + iface=iface, + store=False) thread = threading.Thread(target=sniff_blocking, daemon=True) thread.start() return thread +# ----------------------------- +# PUBLIC API +# ----------------------------- async def start_sniffing(interfaces: List[str]): + global event_loop + event_loop = asyncio.get_running_loop() # <-- IMPORTANT FIX + for iface in interfaces: if iface in sniffer_threads: continue @@ -81,6 +100,6 @@ async def start_sniffing(interfaces: List[str]): async def stop_sniffing(): - # Scapy cannot easily stop sniff(), so we simply kill threads + # Scapy cannot stop sniff(), but threads die on shutdown sniffer_threads.clear() return True