Spaces:
Running
Running
| 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 | |
| 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) | |
| 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() | |