# src/main.py import asyncio import logging import os from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware 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 from src.api import packet_api import src.api.nftables_api as nftables_api # ---- 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( root_path="/api", title="MITM Webserver Backend", description="Backend API for MITM Webserver", version="1.0.0", docs_url="/docs", redoc_url="/redoc", openapi_url="/openapi.json", ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], 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 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") @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 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 # --------------------- # Basic Endpoints # --------------------- @app.get("/hello") def hello(): return {"message": "Hello from FastAPI 🎉"} @app.get("/versions") def versions(): 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"])