"""Rate limiting service using Redis sliding window."""

import time
from dataclasses import dataclass

import structlog
from redis.asyncio import Redis

logger = structlog.get_logger(__name__)


@dataclass
class RateLimitResult:
    """Result of a rate limit check.

    Attributes:
        allowed: Whether the request is allowed.
        remaining: Number of remaining requests in the window.
        reset_at: Unix timestamp when the window resets.
        retry_after: Seconds until the rate limit resets (only if not allowed).
    """

    allowed: bool
    remaining: int
    reset_at: int
    retry_after: int | None = None


async def check_rate_limit(
    redis_client: Redis,
    key_hash: str,
    rate_limit: int,
    window_seconds: int = 60,
) -> RateLimitResult:
    """Check if a request is within rate limits using sliding window.

    Uses Redis INCR with EXPIRE for atomic rate limiting.

    Args:
        redis_client: Async Redis client.
        key_hash: The hashed API key (used as identifier).
        rate_limit: Maximum requests allowed per window.
        window_seconds: Window duration in seconds (default: 60).

    Returns:
        RateLimitResult with allowed status and metadata.
    """
    redis_key = f"rate_limit:{key_hash}:requests"
    current_time = int(time.time())
    window_start = current_time - (current_time % window_seconds)
    reset_at = window_start + window_seconds

    # Get current count
    current_count = await redis_client.get(redis_key)

    if current_count is None:
        # First request in window - initialize counter
        await redis_client.setex(redis_key, window_seconds, 1)
        remaining = rate_limit - 1

        logger.debug(
            "rate_limit_init",
            key_hash_prefix=key_hash[:8],
            remaining=remaining,
            reset_at=reset_at,
        )

        return RateLimitResult(
            allowed=True,
            remaining=remaining,
            reset_at=reset_at,
        )

    count = int(current_count)

    if count >= rate_limit:
        # Rate limit exceeded
        ttl = await redis_client.ttl(redis_key)
        retry_after = max(ttl, 1)  # At least 1 second

        logger.warning(
            "rate_limit_exceeded",
            key_hash_prefix=key_hash[:8],
            count=count,
            limit=rate_limit,
            retry_after=retry_after,
        )

        return RateLimitResult(
            allowed=False,
            remaining=0,
            reset_at=reset_at,
            retry_after=retry_after,
        )

    # Increment counter
    new_count = await redis_client.incr(redis_key)
    remaining = max(rate_limit - new_count, 0)

    logger.debug(
        "rate_limit_check",
        key_hash_prefix=key_hash[:8],
        count=new_count,
        limit=rate_limit,
        remaining=remaining,
    )

    return RateLimitResult(
        allowed=True,
        remaining=remaining,
        reset_at=reset_at,
    )
