import time
from collections import defaultdict, deque
from collections.abc import Awaitable, Callable

from fastapi import Request, Response
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware

from app.auth.jwt import TokenError, decode_token
from app.config import Settings


class RateLimitMiddleware(BaseHTTPMiddleware):
    def __init__(self, app: object, settings: Settings) -> None:
        super().__init__(app)
        self.settings = settings
        self.requests: dict[str, deque[float]] = defaultdict(deque)

    async def dispatch(self, request: Request, call_next: Callable[[Request], Awaitable[Response]]) -> Response:
        if request.url.path == "/health":
            return await call_next(request)
        key = self._key_for_request(request)
        now = time.monotonic()
        window_start = now - self.settings.rate_limit_window_seconds
        bucket = self.requests[key]
        while bucket and bucket[0] < window_start:
            bucket.popleft()
        if len(bucket) >= self.settings.rate_limit_requests:
            return JSONResponse(status_code=429, content={"error": "rate_limited", "message": "Too many requests"})
        bucket.append(now)
        return await call_next(request)

    def _key_for_request(self, request: Request) -> str:
        authorization = request.headers.get("authorization", "")
        if authorization.lower().startswith("bearer "):
            token = authorization.split(" ", 1)[1]
            try:
                payload = decode_token(token, self.settings, "access")
                return f"user:{payload['sub']}"
            except TokenError:
                return f"token:{token[-12:]}"
        host = request.client.host if request.client else "unknown"
        return f"ip:{host}"
