mirror of
https://github.com/agessaman/meshcore-bot.git
synced 2026-08-26 20:40:22 +00:00
feat: per-channel rate limiting
Add ChannelRateLimiter to rate_limiter.py. Configure per-channel cooldowns via [Rate_Limits] channel.<name>_seconds. Integrated into _check_rate_limits() and send_channel_message() in command_manager.py. GET /api/stats/rate_limiters exposes live stats for all four limiter types.
This commit is contained in:
+60
-25
@@ -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.<name>_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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user