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

1import time 

2from typing import Optional 

3 

4from open_webui.env import REDIS_KEY_PREFIX 

5from redis.asyncio import Redis 

6 

7 

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 """ 

13 

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]] = {} 

34 

35 def _bucket_key(self, key: str, bucket_index: int) -> str: 

36 return f'{REDIS_KEY_PREFIX}:ratelimit:{key.lower()}:{bucket_index}' 

37 

38 def _current_bucket(self) -> int: 

39 return int(time.time()) // self.bucket_size 

40 

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] 

46 

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 

54 

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) 

62 

63 async def get_count(self, redis: Redis | None, key: str) -> int: 

64 if not self.enabled: 

65 return 0 

66 

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) 

74 

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) 

78 

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) 

82 

83 attempts = await redis.incr(bucket_key) 

84 if attempts == 1: 

85 await redis.expire(bucket_key, self.window + self.bucket_size) 

86 

87 # Collect buckets 

88 buckets = [self._bucket_key(key, now_bucket - i) for i in range(self.num_buckets + 1)] 

89 

90 counts = await redis.mget(buckets) 

91 total = sum(int(c) for c in counts if c) 

92 

93 return total > self.limit 

94 

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) 

100 

101 def _is_limited_memory(self, key: str) -> bool: 

102 now_bucket = self._current_bucket() 

103 self._prune_memory_store(now_bucket) 

104 

105 current_bucket_counts = self._memory_store.setdefault(now_bucket, {}) 

106 current_bucket_counts[key] = current_bucket_counts.get(key, 0) + 1 

107 

108 total = sum(bucket_counts.get(key, 0) for bucket_counts in self._memory_store.values()) 

109 return total > self.limit 

110 

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())