All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
134 lines
4.0 KiB
Python
134 lines
4.0 KiB
Python
# src/main.py
|
|
import asyncio
|
|
import logging
|
|
import os
|
|
from fastapi import FastAPI
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
|
|
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
|
|
import src.api.nftables_api as nftables_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",
|
|
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"]) |