File size: 4,685 Bytes
732b14f
3c31a2a
 
 
 
 
 
 
732b14f
3c31a2a
 
 
732b14f
 
3c31a2a
 
732b14f
 
 
 
 
 
3c31a2a
 
 
 
732b14f
3c31a2a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
732b14f
 
3c31a2a
732b14f
 
 
3c31a2a
 
732b14f
 
 
 
 
3c31a2a
732b14f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3c31a2a
732b14f
 
 
 
 
 
 
 
3c31a2a
732b14f
3c31a2a
 
 
 
 
 
 
 
 
 
 
 
 
732b14f
3c31a2a
 
 
 
 
 
 
 
 
 
 
732b14f
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
"""Per-tenant sliding-window rate limiter (in-process or Redis-backed)."""

from __future__ import annotations

import logging
import os
import time
from collections import deque
from collections.abc import Awaitable, Callable

from fastapi import HTTPException, Request

from app.config import settings

logger = logging.getLogger(__name__)

_GENERATE_RPM: int = int(
    os.environ.get("RATE_LIMIT_GENERATE_RPM", str(settings.rate_limit_generate_rpm))
)
_READ_RPM: int = int(
    os.environ.get("RATE_LIMIT_READ_RPM", str(settings.rate_limit_read_rpm))
)
_WINDOW_SECS: float = 60.0


class _SlidingWindowLimiter:
    """In-process per-key sliding-window rate limiter."""

    def __init__(self, max_requests: int, window_secs: float) -> None:
        self._max = max_requests
        self._window = window_secs
        self._buckets: dict[str, deque[float]] = {}

    def is_allowed(self, key: str) -> bool:
        now = time.monotonic()
        cutoff = now - self._window
        bucket = self._buckets.setdefault(key, deque())
        while bucket and bucket[0] < cutoff:
            bucket.popleft()
        if len(bucket) >= self._max:
            return False
        bucket.append(now)
        return True

    def reset(self, key: str) -> None:
        self._buckets.pop(key, None)


_generate_mem = _SlidingWindowLimiter(_GENERATE_RPM, _WINDOW_SECS)
_read_mem = _SlidingWindowLimiter(_READ_RPM, _WINDOW_SECS)

_redis_generate = None
_redis_read = None
_use_redis_limits: bool | None = None


async def _ensure_redis_limiters() -> bool:
    global _redis_generate, _redis_read, _use_redis_limits
    if _use_redis_limits is not None:
        return _use_redis_limits
    from app.redis_client import redis_configured

    if not redis_configured():
        _use_redis_limits = False
        return False
    try:
        from app.rate_limit.redis_limiter import RedisSlidingWindowLimiter
        from app.redis_client import get_redis

        client = await get_redis()
        _redis_generate = RedisSlidingWindowLimiter(
            client,
            key_prefix="rics:rl:generate",
            max_requests=_GENERATE_RPM,
            window_secs=_WINDOW_SECS,
        )
        _redis_read = RedisSlidingWindowLimiter(
            client,
            key_prefix="rics:rl:read",
            max_requests=_READ_RPM,
            window_secs=_WINDOW_SECS,
        )
        _use_redis_limits = True
        logger.info("Rate limits using Redis (shared across replicas)")
    except Exception as exc:  # noqa: BLE001
        logger.warning("Redis rate limits unavailable, using in-process: %s", exc)
        _redis_generate = None
        _redis_read = None
        _use_redis_limits = False
    return _use_redis_limits


def reset_rate_limit_backend_for_tests() -> None:
    """Clear cached backend selection (tests only)."""
    global _use_redis_limits, _redis_generate, _redis_read
    _use_redis_limits = None
    _redis_generate = None
    _redis_read = None
    _generate_mem.reset("test")
    _read_mem.reset("test")


async def _is_allowed_generate(tenant_id: str) -> bool:
    if await _ensure_redis_limiters() and _redis_generate is not None:
        return await _redis_generate.is_allowed(tenant_id)
    return _generate_mem.is_allowed(tenant_id)


async def _is_allowed_read(tenant_id: str) -> bool:
    if await _ensure_redis_limiters() and _redis_read is not None:
        return await _redis_read.is_allowed(tenant_id)
    return _read_mem.is_allowed(tenant_id)


async def check_generate(request: Request) -> None:
    tenant_id: str = getattr(request.state, "tenant_id", "anonymous")
    if not await _is_allowed_generate(tenant_id):
        logger.warning("Rate limit exceeded (generate) for tenant=%s", tenant_id)
        raise HTTPException(
            status_code=429,
            detail=(
                f"Rate limit exceeded: at most {_GENERATE_RPM} generation requests "
                f"per minute per tenant. Please wait and retry."
            ),
            headers={"Retry-After": "60"},
        )


async def check_read(request: Request) -> None:
    tenant_id: str = getattr(request.state, "tenant_id", "anonymous")
    if not await _is_allowed_read(tenant_id):
        logger.warning("Rate limit exceeded (read) for tenant=%s", tenant_id)
        raise HTTPException(
            status_code=429,
            detail=(
                f"Rate limit exceeded: at most {_READ_RPM} read requests "
                f"per minute per tenant. Please wait and retry."
            ),
            headers={"Retry-After": "60"},
        )


def rate_limit_backend_label() -> str:
    if _use_redis_limits:
        return "redis"
    return "memory"