Add optional API key auth and rate limiting for REST API

- api_key param: requires 'Authorization: Bearer <key>' header
- rate_limit_rpm param: per-IP sliding-window rate limit (requests/min)
- Both are off by default (backward compatible)
- CLI flags: --api_key, --rate_limit_rpm
- Added 5 unit tests for auth and rate limiting
This commit is contained in:
Aaron Boxer
2026-04-17 09:37:06 -04:00
committed by Aaron Boxer
parent e86d98dd80
commit b648bcb2a4
3 changed files with 133 additions and 1 deletions
+32 -1
View File
@@ -1,6 +1,7 @@
import os
import time
import threading
import collections
import queue
import json
import functools
@@ -8,7 +9,7 @@ import logging
import shutil
import tempfile
from typing import Optional, List
from fastapi import FastAPI, UploadFile, Form
from fastapi import FastAPI, UploadFile, Form, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from starlette.responses import PlainTextResponse, JSONResponse, StreamingResponse
@@ -533,6 +534,8 @@ class TranscriptionServer:
batch_window_ms=50,
raw_pcm_input=False,
metrics_port: int = 0,
api_key: Optional[str] = None,
rate_limit_rpm: int = 0,
segment_post_processor=None):
"""
Run the transcription server.
@@ -612,6 +615,34 @@ class TranscriptionServer:
allow_headers=["*"], # Allows all headers
)
# Optional API key authentication
if api_key:
@app.middleware("http")
async def _check_api_key(request: Request, call_next):
auth = request.headers.get("Authorization", "")
if auth != f"Bearer {api_key}":
return JSONResponse({"error": "Invalid or missing API key"}, status_code=401)
return await call_next(request)
# Optional rate limiting (requests per minute per client IP)
if rate_limit_rpm > 0:
_rate_lock = threading.Lock()
_rate_buckets: dict = {} # ip -> deque of timestamps
@app.middleware("http")
async def _rate_limit(request: Request, call_next):
client_ip = request.client.host if request.client else "unknown"
now = time.time()
with _rate_lock:
bucket = _rate_buckets.setdefault(client_ip, collections.deque())
# Discard entries older than 60s
while bucket and bucket[0] < now - 60:
bucket.popleft()
if len(bucket) >= rate_limit_rpm:
return JSONResponse({"error": "Rate limit exceeded"}, status_code=429)
bucket.append(now)
return await call_next(request)
@app.post("/v1/audio/transcriptions")
async def transcribe(