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