MediaRouter / app /security /rate_limit.py
basyx's picture
Upload 236 files
e1104b3 verified
Raw
History Blame Contribute Delete
6.07 kB
from __future__ import annotations
import asyncio
import math
import time
from collections import defaultdict, deque
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
from app.security.context import AuthContext
from app.security.database import SecurityDatabase
from app.security.errors import RateLimitError
from app.security.models import RateLimit, utcnow
@dataclass(slots=True)
class RateLimitLease:
limiter: APIKeyRateLimiter
api_key_id: str
concurrent: bool
async def release(self) -> None:
if self.concurrent:
await self.limiter.release_job(self.api_key_id)
class APIKeyRateLimiter:
"""Low-latency per-key windows with durable aggregate counters for auditing."""
def __init__(self, database: SecurityDatabase) -> None:
self.database = database
self._lock = asyncio.Lock()
self._requests: dict[str, deque[float]] = defaultdict(deque)
self._uploads: dict[str, deque[float]] = defaultdict(deque)
self._concurrent: dict[str, int] = defaultdict(int)
self._daily_bytes: dict[tuple[str, str], int] = defaultdict(int)
self._categories: dict[tuple[str, str], deque[float]] = defaultdict(deque)
async def acquire(
self,
context: AuthContext,
*,
is_job: bool,
is_upload: bool,
uploaded_bytes: int,
) -> RateLimitLease:
now = time.time()
today = datetime.now(timezone.utc).date().isoformat()
retry_after = 0
async with self._lock:
requests = self._requests[context.api_key_id]
self._prune(requests, now - 60)
if len(requests) >= context.requests_per_minute:
retry_after = max(1, math.ceil(requests[0] + 60 - now))
uploads = self._uploads[context.api_key_id]
self._prune(uploads, now - 3600)
if not retry_after and is_upload and len(uploads) >= context.uploads_per_hour:
retry_after = max(1, math.ceil(uploads[0] + 3600 - now))
daily_key = (context.api_key_id, today)
daily_total = self._daily_bytes[daily_key]
if (
not retry_after
and uploaded_bytes
and daily_total + uploaded_bytes > context.processing_bytes_per_day
):
tomorrow = datetime.now(timezone.utc).replace(
hour=0, minute=0, second=0, microsecond=0
) + timedelta(days=1)
retry_after = max(1, int((tomorrow - datetime.now(timezone.utc)).total_seconds()))
if (
not retry_after
and is_job
and self._concurrent[context.api_key_id] >= context.concurrent_jobs
):
retry_after = 1
if retry_after:
raise RateLimitError(retry_after)
requests.append(now)
if is_upload:
uploads.append(now)
if uploaded_bytes:
self._daily_bytes[daily_key] += uploaded_bytes
if is_job:
self._concurrent[context.api_key_id] += 1
await self._record(context.api_key_id, "requests_minute", 1, 0, 60)
if is_upload:
await self._record(context.api_key_id, "uploads_hour", 1, 0, 3600)
if uploaded_bytes:
await self._record(
context.api_key_id, "processing_bytes_day", 0, uploaded_bytes, 86_400
)
return RateLimitLease(self, context.api_key_id, is_job)
async def release_job(self, api_key_id: str) -> None:
async with self._lock:
self._concurrent[api_key_id] = max(0, self._concurrent[api_key_id] - 1)
async def acquire_category(
self,
context: AuthContext,
category: str,
*,
limit: int,
window_seconds: int,
) -> None:
"""Reserve an independent social-operation bucket for a key.
Generic API limits still apply in middleware. These smaller buckets
prevent OAuth, publishing, scheduling, and analytics traffic from
starving each other when the social subsystem is enabled.
"""
now = time.time()
key = (context.api_key_id, category)
async with self._lock:
values = self._categories[key]
self._prune(values, now - window_seconds)
if len(values) >= limit:
raise RateLimitError(max(1, math.ceil(values[0] + window_seconds - now)))
values.append(now)
await self._record(context.api_key_id, category, 1, 0, window_seconds)
@staticmethod
def _prune(values: deque[float], cutoff: float) -> None:
while values and values[0] <= cutoff:
values.popleft()
async def _record(
self, api_key_id: str, bucket_type: str, count: int, units: int, seconds: int
) -> None:
now = datetime.now(timezone.utc)
epoch = int(now.timestamp())
bucket_start = datetime.fromtimestamp(epoch - (epoch % seconds), timezone.utc)
async with self.database.session() as session:
statement = sqlite_insert(RateLimit).values(
api_key_id=api_key_id,
bucket_type=bucket_type,
bucket_start=bucket_start,
count=count,
units=units,
updated_at=utcnow(),
)
statement = statement.on_conflict_do_update(
index_elements=[
RateLimit.api_key_id,
RateLimit.bucket_type,
RateLimit.bucket_start,
],
set_={
"count": RateLimit.count + statement.excluded.count,
"units": RateLimit.units + statement.excluded.units,
"updated_at": utcnow(),
},
)
await session.execute(statement)
await session.commit()