import time
import logging
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from starlette.responses import Response
from starlette.middleware.cors import CORSMiddleware
from src.core.config import settings

logger = logging.getLogger(__name__)


class LoggingMiddleware(BaseHTTPMiddleware):
    async def dispatch(self, request: Request, call_next):
        start_time = time.time()
        response = await call_next(request)
        process_time = time.time() - start_time

        logger.info(
            f"{request.method} {request.url.path} completed in {process_time:.2f} seconds"
        )

        return response


def custom_middleware(app):
    # Add logging middleware
    app.add_middleware(LoggingMiddleware)

    # Add CORS middleware
    app.add_middleware(
        CORSMiddleware,
        allow_origins=settings.ORIGINS,
        allow_credentials=True,
        allow_methods=["*"],
        allow_headers=["*"],
    )

    # Add other middlewares as needed
    # Example: Add a custom header to each response
    @app.middleware("http")
    async def add_custom_header(request: Request, call_next):
        response = await call_next(request)
        response.headers["X-Custom-Header"] = "MyCustomHeaderValue"
        return response

    return app
