diff --git a/modules/rate_limiter.py b/modules/rate_limiter.py index 29b9c60..321d82a 100644 --- a/modules/rate_limiter.py +++ b/modules/rate_limiter.py @@ -4,9 +4,9 @@ Rate limiting functionality for the MeshCore Bot Controls how often messages can be sent to prevent spam """ -import time import asyncio -from typing import Optional, Dict, List +import time +from typing import Optional class PerUserRateLimiter: @@ -20,8 +20,8 @@ class PerUserRateLimiter: def __init__(self, seconds: float, max_entries: int = 1000): self.seconds = seconds self.max_entries = max_entries - self._last_send: Dict[str, float] = {} - self._order: List[str] = [] # keys in insertion order for oldest-first eviction + self._last_send: dict[str, float] = {} + self._order: list[str] = [] # keys in insertion order for oldest-first eviction def _evict_if_needed(self, new_key: str) -> None: """Evict oldest entry if at capacity and new_key is not already present.""" @@ -59,30 +59,30 @@ class PerUserRateLimiter: class RateLimiter: """Rate limiting for message sending""" - + def __init__(self, seconds: int): self.seconds = seconds self.last_send = 0 self._total_sends = 0 self._total_throttled = 0 - + def can_send(self) -> bool: """Check if we can send a message""" can = time.time() - self.last_send >= self.seconds if not can: self._total_throttled += 1 return can - + def time_until_next(self) -> float: """Get time until next allowed send""" elapsed = time.time() - self.last_send return max(0, self.seconds - elapsed) - + def record_send(self): """Record that we sent a message""" self.last_send = time.time() self._total_sends += 1 - + def get_stats(self) -> dict: """Get rate limiter statistics""" total_attempts = self._total_sends + self._total_throttled @@ -96,37 +96,37 @@ class RateLimiter: class BotTxRateLimiter: """Rate limiting for bot transmission to prevent network overload""" - + def __init__(self, seconds: float = 1.0): self.seconds = seconds self.last_tx = 0 self._total_tx = 0 self._total_throttled = 0 - + def can_tx(self) -> bool: """Check if bot can transmit a message""" can = time.time() - self.last_tx >= self.seconds if not can: self._total_throttled += 1 return can - + def time_until_next_tx(self) -> float: """Get time until next allowed transmission""" elapsed = time.time() - self.last_tx return max(0, self.seconds - elapsed) - + def record_tx(self): """Record that bot transmitted a message""" self.last_tx = time.time() self._total_tx += 1 - + async def wait_for_tx(self): """Wait until bot can transmit (async)""" while not self.can_tx(): wait_time = self.time_until_next_tx() if wait_time > 0: await asyncio.sleep(wait_time + 0.05) # Small buffer - + def get_stats(self) -> dict: """Get rate limiter statistics""" total_attempts = self._total_tx + self._total_throttled @@ -138,50 +138,85 @@ class BotTxRateLimiter: } +class ChannelRateLimiter: + """Per-channel rate limiting: minimum seconds between bot messages on the same channel. + + Channel names are mapped to ``RateLimiter`` instances using limits loaded from + the ``[Rate_Limits]`` config section (``channel._seconds = N``). + Channels without an explicit limit are unrestricted. + """ + + def __init__(self, channel_limits: dict[str, float]): + self._limiters: dict[str, RateLimiter] = { + channel: RateLimiter(int(max(1, seconds))) + for channel, seconds in channel_limits.items() + if seconds > 0 + } + + def can_send(self, channel: str) -> bool: + limiter = self._limiters.get(channel) + return limiter.can_send() if limiter else True + + def time_until_next(self, channel: str) -> float: + limiter = self._limiters.get(channel) + return limiter.time_until_next() if limiter else 0.0 + + def record_send(self, channel: str) -> None: + limiter = self._limiters.get(channel) + if limiter: + limiter.record_send() + + def get_stats(self) -> dict[str, dict]: + return {ch: lim.get_stats() for ch, lim in self._limiters.items()} + + def channels(self) -> list[str]: + return list(self._limiters.keys()) + + class NominatimRateLimiter: """Rate limiting for Nominatim geocoding API requests - + Nominatim policy: Maximum 1 request per second We'll be conservative and use 1.1 seconds to ensure compliance """ - + def __init__(self, seconds: float = 1.1): self.seconds = seconds self.last_request = 0 self._lock: Optional[asyncio.Lock] = None self._total_requests = 0 self._total_throttled = 0 - + def _get_lock(self) -> asyncio.Lock: """Lazily initialize the async lock""" if self._lock is None: self._lock = asyncio.Lock() return self._lock - + def can_request(self) -> bool: """Check if we can make a Nominatim request""" can = time.time() - self.last_request >= self.seconds if not can: self._total_throttled += 1 return can - + def time_until_next(self) -> float: """Get time until next allowed request""" elapsed = time.time() - self.last_request return max(0, self.seconds - elapsed) - + def record_request(self): """Record that we made a Nominatim request""" self.last_request = time.time() self._total_requests += 1 - + async def wait_for_request(self): """Wait until we can make a Nominatim request (async)""" while not self.can_request(): wait_time = self.time_until_next() if wait_time > 0: await asyncio.sleep(wait_time + 0.05) # Small buffer - + async def wait_and_request(self) -> None: """Wait until a request can be made, then mark request time (thread-safe)""" async with self._get_lock(): @@ -191,14 +226,14 @@ class NominatimRateLimiter: await asyncio.sleep(self.seconds - time_since_last) self.last_request = time.time() self._total_requests += 1 - + def wait_for_request_sync(self): """Wait until we can make a Nominatim request (synchronous)""" while not self.can_request(): wait_time = self.time_until_next() if wait_time > 0: time.sleep(wait_time + 0.05) # Small buffer - + def get_stats(self) -> dict: """Get rate limiter statistics""" total_attempts = self._total_requests + self._total_throttled diff --git a/tests/test_rate_limiter.py b/tests/test_rate_limiter.py index 4193fb5..6145f81 100644 --- a/tests/test_rate_limiter.py +++ b/tests/test_rate_limiter.py @@ -1,11 +1,8 @@ """Tests for modules.rate_limiter.""" import time -from unittest.mock import patch -import pytest - -from modules.rate_limiter import RateLimiter, PerUserRateLimiter +from modules.rate_limiter import PerUserRateLimiter, RateLimiter class TestRateLimiter: