file structure and comments unified
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s
This commit is contained in:
@@ -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"])
|
||||
|
||||
Reference in New Issue
Block a user