#!/usr/bin/env python3
"""Rate limiter unit tests (Phase 2 — 2d)

Tests the Redis-backed sliding window rate limiter:
- Allow under limit
- Deny over limit
- Window expiry resets counter
- Redis unavailable → fail-open
- get_remaining() accuracy
- flush() clears keys
"""
import os
import sys
import time

import pytest

sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))

from app.services.ratelimit import RateLimiter


class FakeRedis:
    """Minimal Redis mock for unit testing without a real Redis instance."""

    def __init__(self):
        self._store = {}
        self._ttls = {}

    def incr(self, key):
        # Check TTL expiry before incrementing
        if key in self._ttls and time.time() > self._ttls[key]:
            del self._store[key]
            del self._ttls[key]
        val = self._store.get(key, 0) + 1
        self._store[key] = val
        return val

    def expire(self, key, ttl):
        self._ttls[key] = time.time() + ttl
        return True

    def get(self, key):
        if key in self._ttls and time.time() > self._ttls[key]:
            del self._store[key]
            del self._ttls[key]
            return None
        return str(self._store.get(key, "")) or None

    def ttl(self, key):
        if key in self._ttls:
            remaining = self._ttls[key] - time.time()
            return max(0, int(remaining))
        return -2

    def delete(self, *keys):
        for key in keys:
            self._store.pop(key, None)
            self._ttls.pop(key, None)

    def keys(self, pattern="*"):
        import fnmatch
        return [k for k in self._store if fnmatch.fnmatch(k, pattern)]

    def pipeline(self):
        return _FakePipeline(self)


class _FakePipeline:
    def __init__(self, redis):
        self._redis = redis
        self._ops = []

    def incr(self, key):
        self._ops.append(("incr", key))
        return self

    def expire(self, key, ttl):
        self._ops.append(("expire", key, ttl))
        return self

    def execute(self):
        results = []
        for op in self._ops:
            if op[0] == "incr":
                results.append(self._redis.incr(op[1]))
            elif op[0] == "expire":
                results.append(self._redis.expire(op[1], op[2]))
        self._ops = []
        return results


class TestRateLimiterAllow:
    """Test basic allow/deny logic."""

    @pytest.fixture
    def limiter(self):
        r = FakeRedis()
        lim = RateLimiter(r)
        return lim

    def test_allow_under_limit(self, limiter):
        result = limiter.allow("test:allow", max_requests=5, window=60)
        assert result is True

    def test_deny_over_limit(self, limiter):
        # Fill up to the limit
        for _ in range(5):
            limiter.allow("test:deny", max_requests=5, window=60)
        # 6th request should be denied
        result = limiter.allow("test:deny", max_requests=5, window=60)
        assert result is False

    def test_exact_limit_allowed(self, limiter):
        # Exactly at the limit should still be allowed
        for i in range(3):
            result = limiter.allow("test:exact", max_requests=3, window=60)
        assert result is True
        # One over
        result = limiter.allow("test:exact", max_requests=3, window=60)
        assert result is False


class TestRateLimiterCheck:
    """Test check() return values."""

    @pytest.fixture
    def limiter(self):
        r = FakeRedis()
        lim = RateLimiter(r)
        return lim

    def test_check_returns_tuple(self, limiter):
        allowed, remaining, reset_at = limiter.check("test:check", max_requests=10, window=60)
        assert isinstance(allowed, bool)
        assert isinstance(remaining, int)
        assert isinstance(reset_at, float)

    def test_check_remaining_decrements(self, limiter):
        _, remaining, _ = limiter.check("test:rem", max_requests=5, window=60)
        assert remaining == 4

        for _ in range(4):
            limiter.check("test:rem", max_requests=5, window=60)

        _, remaining, _ = limiter.check("test:rem", max_requests=5, window=60)
        assert remaining == 0

    def test_check_reset_at_future(self, limiter):
        _, _, reset_at = limiter.check("test:reset", max_requests=10, window=60)
        assert reset_at > time.time()
        assert reset_at <= time.time() + 61


class TestRateLimiterWindowExpiry:
    """Test that counters reset after window expires."""

    @pytest.fixture
    def limiter(self):
        r = FakeRedis()
        lim = RateLimiter(r)
        return lim

    def test_window_expiry_resets(self, limiter):
        # Fill up with a 1-second window
        for _ in range(5):
            limiter.allow("test:expiry", max_requests=5, window=1)
        assert limiter.allow("test:expiry", max_requests=5, window=1) is False

        # Wait for expiry
        time.sleep(1.1)

        # Should be allowed again
        assert limiter.allow("test:expiry", max_requests=5, window=1) is True


class TestRateLimiterFailOpen:
    """Test fail-open behavior when Redis is unavailable."""

    def test_no_redis_allow(self):
        lim = RateLimiter()  # No redis_client
        result = lim.allow("test:noredis")
        assert result is True

    def test_no_redis_check(self):
        lim = RateLimiter()
        allowed, remaining, reset_at = lim.check("test:noredis", max_requests=10, window=60)
        assert allowed is True
        assert remaining == 10
        assert reset_at == 0

    def test_no_redis_get_remaining(self):
        lim = RateLimiter()
        remaining, reset_at = lim.get_remaining("test:noredis", max_requests=10, window=60)
        assert remaining == 10
        assert reset_at == 0

    def test_no_redis_flush(self):
        lim = RateLimiter()
        # Should not raise
        lim.flush("test:*")


class TestRateLimiterGetRemaining:
    """Test get_remaining() accuracy."""

    @pytest.fixture
    def limiter(self):
        r = FakeRedis()
        lim = RateLimiter(r)
        return lim

    def test_get_remaining_fresh(self, limiter):
        remaining, _ = limiter.get_remaining("test:fresh", max_requests=10, window=60)
        assert remaining == 10

    def test_get_remaining_after_usage(self, limiter):
        limiter.allow("test:used", max_requests=10, window=60)
        limiter.allow("test:used", max_requests=10, window=60)
        remaining, _ = limiter.get_remaining("test:used", max_requests=10, window=60)
        assert remaining == 8

    def test_get_remaining_zero(self, limiter):
        for _ in range(5):
            limiter.allow("test:zero", max_requests=5, window=60)
        remaining, _ = limiter.get_remaining("test:zero", max_requests=5, window=60)
        assert remaining == 0


class TestRateLimiterFlush:
    """Test flush() clears keys."""

    @pytest.fixture
    def limiter(self):
        r = FakeRedis()
        lim = RateLimiter(r)
        return lim

    def test_flush_clears_keys(self, limiter):
        limiter.allow("flush:test:a", max_requests=10, window=60)
        limiter.allow("flush:test:b", max_requests=10, window=60)

        limiter.flush("flush:test:*")

        # Keys should be cleared — fresh counts
        remaining, _ = limiter.get_remaining("flush:test:a", max_requests=10, window=60)
        assert remaining == 10

    def test_flush_no_match(self, limiter):
        limiter.allow("flush:other", max_requests=10, window=60)
        limiter.flush("flush:nomatch:*")
        # Should not raise on no matches
        remaining, _ = limiter.get_remaining("flush:other", max_requests=10, window=60)
        assert remaining == 9


if __name__ == "__main__":
    pytest.main([__file__, "-v", "--tb=short"])