lumaspeech-backend / api /middleware /security_middleware.py
yasmine hemmati
Initial deployment: LumaSpeech Backend API with GPU support
d646f8a
Raw
History Blame Contribute Delete
6.83 kB
"""
Security middleware for production environment
"""
import time
import hashlib
from collections import defaultdict
from datetime import datetime, timedelta
from typing import Dict, Optional
from fastapi import Request, HTTPException, status
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.trustedhost import TrustedHostMiddleware
from starlette.middleware.base import BaseHTTPMiddleware
from config.security import SecurityConfig
class RateLimitMiddleware(BaseHTTPMiddleware):
"""Rate limiting middleware to prevent brute force attacks"""
def __init__(self, app, calls: int = 100, period: int = 60):
super().__init__(app)
self.calls = calls
self.period = period
self.clients: Dict[str, list] = defaultdict(list)
self.blocked_clients: Dict[str, float] = {} # IP -> blocked_until timestamp
# Specific limits for sensitive endpoints
self.endpoint_limits = {
'/api/auth/patient/login': (5, 900), # 5 attempts per 15 minutes
'/api/auth/therapist/login': (5, 900),
'/api/auth/patient/signup': (3, 3600), # 3 signups per hour
'/api/auth/therapist/signup': (3, 3600),
}
def get_client_ip(self, request: Request) -> str:
"""Get real client IP considering proxies"""
# Check for real IP behind proxy
forwarded = request.headers.get("X-Forwarded-For")
if forwarded:
return forwarded.split(",")[0]
return request.client.host
async def dispatch(self, request: Request, call_next):
client_ip = self.get_client_ip(request)
path = request.url.path
current_time = time.time()
# Check if client is blocked
if client_ip in self.blocked_clients:
if current_time < self.blocked_clients[client_ip]:
remaining = int(self.blocked_clients[client_ip] - current_time)
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail=f"Too many requests. Please try again in {remaining} seconds."
)
else:
# Unblock client
del self.blocked_clients[client_ip]
# Get limits for this endpoint
if path in self.endpoint_limits:
limit_calls, limit_period = self.endpoint_limits[path]
else:
limit_calls, limit_period = self.calls, self.period
# Clean old requests
min_time = current_time - limit_period
self.clients[client_ip] = [
req_time for req_time in self.clients[client_ip]
if req_time > min_time
]
# Check rate limit
if len(self.clients[client_ip]) >= limit_calls:
# Block client temporarily
self.blocked_clients[client_ip] = current_time + 300 # Block for 5 minutes
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Too many requests. Please slow down."
)
# Record this request
self.clients[client_ip].append(current_time)
# Process request
response = await call_next(request)
return response
class SecurityHeadersMiddleware(BaseHTTPMiddleware):
"""Add security headers to all responses"""
async def dispatch(self, request: Request, call_next):
response = await call_next(request)
# Security headers
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["X-Frame-Options"] = "DENY"
response.headers["X-XSS-Protection"] = "1; mode=block"
response.headers["Strict-Transport-Security"] = "max-age=31536000; includeSubDomains"
response.headers["Content-Security-Policy"] = "default-src 'self'"
response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin"
response.headers["Permissions-Policy"] = "geolocation=(), microphone=(), camera=()"
return response
class LoginAttemptTracker:
"""Track failed login attempts for additional security"""
def __init__(self):
self.attempts: Dict[str, list] = defaultdict(list)
self.locked_accounts: Dict[str, float] = {}
def record_failed_attempt(self, identifier: str) -> None:
"""Record a failed login attempt"""
current_time = time.time()
self.attempts[identifier].append(current_time)
# Clean old attempts (older than 15 minutes)
min_time = current_time - SecurityConfig.LOGIN_ATTEMPTS_WINDOW
self.attempts[identifier] = [
t for t in self.attempts[identifier] if t > min_time
]
# Lock account if too many attempts
if len(self.attempts[identifier]) >= SecurityConfig.LOGIN_ATTEMPTS_LIMIT:
self.locked_accounts[identifier] = current_time + 900 # Lock for 15 minutes
def is_locked(self, identifier: str) -> tuple[bool, Optional[int]]:
"""Check if account is locked"""
if identifier in self.locked_accounts:
current_time = time.time()
if current_time < self.locked_accounts[identifier]:
remaining = int(self.locked_accounts[identifier] - current_time)
return True, remaining
else:
# Unlock account
del self.locked_accounts[identifier]
if identifier in self.attempts:
del self.attempts[identifier]
return False, None
def clear_attempts(self, identifier: str) -> None:
"""Clear attempts after successful login"""
if identifier in self.attempts:
del self.attempts[identifier]
if identifier in self.locked_accounts:
del self.locked_accounts[identifier]
# Global instance for login tracking
login_tracker = LoginAttemptTracker()
def setup_security_middleware(app):
"""Setup all security middleware for the application"""
# CORS middleware
app.add_middleware(
CORSMiddleware,
allow_origins=SecurityConfig.ALLOWED_ORIGINS,
allow_credentials=True,
allow_methods=["GET", "POST", "PUT", "DELETE", "OPTIONS"],
allow_headers=["*"],
expose_headers=["*"]
)
# Trusted host middleware (prevent host header attacks)
if os.getenv('ENVIRONMENT') == 'production':
allowed_hosts = os.getenv('ALLOWED_HOSTS', 'localhost').split(',')
app.add_middleware(TrustedHostMiddleware, allowed_hosts=allowed_hosts)
# Rate limiting
app.add_middleware(RateLimitMiddleware)
# Security headers
app.add_middleware(SecurityHeadersMiddleware)
return app