Coverage for open_webui/utils/rate_limit.py: 50%
68 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import time
2from typing import Optional
4from open_webui.env import REDIS_KEY_PREFIX
5from redis.asyncio import Redis
8class RateLimiter:
9 """
10 General-purpose rate limiter using Redis with a rolling window strategy.
11 Falls back to in-memory storage if Redis is not available.
12 """
14 def __init__(
15 self,
16 limit: int,
17 window: int,
18 bucket_size: int = 60,
19 enabled: bool = True,
20 ):
21 """
22 :param limit: Max allowed events in the window
23 :param window: Time window in seconds
24 :param bucket_size: Bucket resolution
25 :param enabled: Turn on/off rate limiting globally
26 """
27 self.limit = limit
28 self.window = window
29 self.bucket_size = bucket_size
30 self.num_buckets = window // bucket_size
31 self.enabled = enabled
32 # bucket index -> rate-limit key -> hits
33 self._memory_store: dict[int, dict[str, int]] = {}
35 def _bucket_key(self, key: str, bucket_index: int) -> str:
36 return f'{REDIS_KEY_PREFIX}:ratelimit:{key.lower()}:{bucket_index}'
38 def _current_bucket(self) -> int:
39 return int(time.time()) // self.bucket_size
41 def _prune_memory_store(self, now_bucket: int) -> None:
42 min_bucket = now_bucket - self.num_buckets
43 expired = [bucket_index for bucket_index in self._memory_store if bucket_index < min_bucket]
44 for bucket_index in expired:
45 del self._memory_store[bucket_index]
47 async def is_limited(self, redis: Redis | None, key: str) -> bool:
48 """
49 Main rate-limit check.
50 Gracefully handles missing or failing Redis.
51 """
52 if not self.enabled: 52 ↛ 53line 52 didn't jump to line 53 because the condition on line 52 was never true
53 return False
55 if redis is not None: 55 ↛ 56line 55 didn't jump to line 56 because the condition on line 55 was never true
56 try:
57 return await self._is_limited_redis(redis, key)
58 except Exception:
59 return self._is_limited_memory(key)
60 else:
61 return self._is_limited_memory(key)
63 async def get_count(self, redis: Redis | None, key: str) -> int:
64 if not self.enabled:
65 return 0
67 if redis is not None:
68 try:
69 return await self._get_count_redis(redis, key)
70 except Exception:
71 return self._get_count_memory(key)
72 else:
73 return self._get_count_memory(key)
75 async def remaining(self, redis: Redis | None, key: str) -> int:
76 used = await self.get_count(redis, key)
77 return max(0, self.limit - used)
79 async def _is_limited_redis(self, redis: Redis, key: str) -> bool:
80 now_bucket = self._current_bucket()
81 bucket_key = self._bucket_key(key, now_bucket)
83 attempts = await redis.incr(bucket_key)
84 if attempts == 1:
85 await redis.expire(bucket_key, self.window + self.bucket_size)
87 # Collect buckets
88 buckets = [self._bucket_key(key, now_bucket - i) for i in range(self.num_buckets + 1)]
90 counts = await redis.mget(buckets)
91 total = sum(int(c) for c in counts if c)
93 return total > self.limit
95 async def _get_count_redis(self, redis: Redis, key: str) -> int:
96 now_bucket = self._current_bucket()
97 buckets = [self._bucket_key(key, now_bucket - i) for i in range(self.num_buckets + 1)]
98 counts = await redis.mget(buckets)
99 return sum(int(c) for c in counts if c)
101 def _is_limited_memory(self, key: str) -> bool:
102 now_bucket = self._current_bucket()
103 self._prune_memory_store(now_bucket)
105 current_bucket_counts = self._memory_store.setdefault(now_bucket, {})
106 current_bucket_counts[key] = current_bucket_counts.get(key, 0) + 1
108 total = sum(bucket_counts.get(key, 0) for bucket_counts in self._memory_store.values())
109 return total > self.limit
111 def _get_count_memory(self, key: str) -> int:
112 now_bucket = self._current_bucket()
113 self._prune_memory_store(now_bucket)
114 return sum(bucket_counts.get(key, 0) for bucket_counts in self._memory_store.values())