| """ |
| Security Middleware |
| Provides rate limiting, input validation, and security headers |
| """ |
|
|
| from datetime import datetime, timedelta |
| import hashlib |
| import logging |
| import re |
| import secrets |
| import time |
| from typing import Any, Dict, Optional |
| from fastapi import HTTPException, Request, Response |
| from starlette.middleware.base import BaseHTTPMiddleware |
| from starlette.responses import JSONResponse |
|
|
| from core.auth import get_password_hash as secure_hash_password |
|
|
| |
| security_logger = logging.getLogger("atom.security") |
|
|
|
|
| class RateLimitMiddleware(BaseHTTPMiddleware): |
| """Rate limiting middleware with configurable limits""" |
|
|
| def __init__(self, app, requests_per_minute: int = 60, burst_size: int = 10): |
| super().__init__(app) |
| self.requests_per_minute = requests_per_minute |
| self.burst_size = burst_size |
| self.clients: Dict[str, Dict[str, Any]] = {} |
|
|
| async def dispatch(self, request: Request, call_next): |
| |
| client_ip = self._get_client_ip(request) |
|
|
| |
| if self._is_rate_limited(client_ip): |
| security_logger.warning( |
| f"Rate limit exceeded for IP: {client_ip} - {request.method} {request.url.path}" |
| ) |
| return JSONResponse( |
| status_code=429, |
| content={ |
| "error": { |
| "type": "rate_limit_exceeded", |
| "message": "Too many requests. Please try again later.", |
| "retry_after": 60 |
| } |
| }, |
| headers={ |
| "Retry-After": "60", |
| "X-RateLimit-Limit": str(self.requests_per_minute), |
| "X-RateLimit-Remaining": "0", |
| "X-RateLimit-Reset": str(int(time.time()) + 60) |
| } |
| ) |
|
|
| |
| response = await call_next(request) |
|
|
| |
| client_data = self.clients.get(client_ip, {}) |
| remaining = max(0, self.requests_per_minute - client_data.get("count", 0)) |
| reset_time = int(client_data.get("reset_time", time.time() + 60)) |
|
|
| response.headers["X-RateLimit-Limit"] = str(self.requests_per_minute) |
| response.headers["X-RateLimit-Remaining"] = str(remaining) |
| response.headers["X-RateLimit-Reset"] = str(reset_time) |
|
|
| return response |
|
|
| def _get_client_ip(self, request: Request) -> str: |
| """Get client IP from request""" |
| |
| forwarded_for = request.headers.get("x-forwarded-for") |
| if forwarded_for: |
| return forwarded_for.split(",")[0].strip() |
|
|
| real_ip = request.headers.get("x-real-ip") |
| if real_ip: |
| return real_ip |
|
|
| return request.client.host if request.client else "unknown" |
|
|
| def _is_rate_limited(self, client_ip: str) -> bool: |
| """Check if client has exceeded rate limit""" |
| current_time = time.time() |
|
|
| |
| if client_ip not in self.clients: |
| self.clients[client_ip] = { |
| "count": 0, |
| "reset_time": current_time + 60, |
| "burst_tokens": self.burst_size |
| } |
|
|
| client_data = self.clients[client_ip] |
|
|
| |
| if current_time > client_data["reset_time"]: |
| client_data["count"] = 0 |
| client_data["reset_time"] = current_time + 60 |
| client_data["burst_tokens"] = self.burst_size |
|
|
| |
| if client_data["burst_tokens"] > 0: |
| client_data["burst_tokens"] -= 1 |
| client_data["count"] += 1 |
| return False |
|
|
| |
| if client_data["count"] >= self.requests_per_minute: |
| return True |
|
|
| |
| client_data["count"] += 1 |
| return False |
|
|
|
|
| class InputValidationMiddleware(BaseHTTPMiddleware): |
| """Input validation middleware for security""" |
|
|
| def __init__(self, app): |
| super().__init__(app) |
| |
| self.malicious_patterns = [ |
| r'<script[^>]*>.*?</script>', |
| r'javascript:', |
| r'on\w+\s*=', |
| r'union\s+select', |
| r'drop\s+table', |
| r'exec\(', |
| r'eval\(', |
| r'system\(', |
| ] |
|
|
| async def dispatch(self, request: Request, call_next): |
| |
| if not self._validate_query_params(request): |
| security_logger.warning( |
| f"Malicious query params detected: {request.query_params}" |
| ) |
| return JSONResponse( |
| status_code=400, |
| content={ |
| "error": { |
| "type": "invalid_input", |
| "message": "Invalid request parameters" |
| } |
| } |
| ) |
|
|
| |
| if request.method in ["POST", "PUT", "PATCH"]: |
| try: |
| |
| body = await request.body() |
| body_str = body.decode('utf-8', errors='ignore') |
|
|
| |
| if not self._validate_content(body_str): |
| security_logger.warning( |
| f"Malicious content detected in body: {body_str[:200]}..." |
| ) |
| return JSONResponse( |
| status_code=400, |
| content={ |
| "error": { |
| "type": "invalid_input", |
| "message": "Invalid request content" |
| } |
| } |
| ) |
|
|
| |
| |
| |
| request._body = body |
|
|
| except Exception as e: |
| logger.warning(f"Could not read request body for security check: {e}") |
| |
|
|
| return await call_next(request) |
|
|
| def _validate_query_params(self, request: Request) -> bool: |
| """Validate query parameters""" |
| for param_name, param_value in request.query_params.items(): |
| |
| if self._contains_malicious_content(str(param_value)): |
| return False |
|
|
| |
| if len(str(param_value)) > 1000: |
| return False |
|
|
| return True |
|
|
| def _validate_content(self, content: str) -> bool: |
| """Validate request content""" |
| |
| if self._contains_malicious_content(content): |
| return False |
|
|
| |
| if len(content) > 10 * 1024 * 1024: |
| return False |
|
|
| return True |
|
|
| def _contains_malicious_content(self, content: str) -> bool: |
| """Check if content contains malicious patterns""" |
| content_lower = content.lower() |
| for pattern in self.malicious_patterns: |
| if re.search(pattern, content_lower, re.IGNORECASE | re.MULTILINE): |
| return True |
| return False |
|
|
|
|
| class SecurityHeadersMiddleware(BaseHTTPMiddleware): |
| """Add security headers to responses""" |
|
|
| async def dispatch(self, request: Request, call_next): |
| response = await call_next(request) |
|
|
| |
| 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["Referrer-Policy"] = "strict-origin-when-cross-origin" |
| response.headers["Content-Security-Policy"] = ( |
| "default-src 'self'; " |
| "script-src 'self' 'unsafe-inline' 'unsafe-eval'; " |
| "style-src 'self' 'unsafe-inline'; " |
| "img-src 'self' data: https:; " |
| "font-src 'self' data:; " |
| "connect-src 'self' ws: wss: https:;" |
| ) |
| response.headers["Permissions-Policy"] = ( |
| "camera=(), microphone=(), geolocation=(), " |
| "payment=(), usb=(), magnetometer=(), gyroscope=()" |
| ) |
|
|
| return response |
|
|
|
|
| class CSRFProtectionMiddleware(BaseHTTPMiddleware): |
| """CSRF protection middleware (simplified for MVP)""" |
|
|
| def __init__(self, app): |
| super().__init__(app) |
| self.csrf_tokens = {} |
| self.token_expiry = 3600 |
|
|
| async def dispatch(self, request: Request, call_next): |
| |
| if request.method in ["GET", "HEAD", "OPTIONS"]: |
| return await call_next(request) |
|
|
| |
| if request.method in ["POST", "PUT", "DELETE", "PATCH"]: |
| csrf_token = request.headers.get("X-CSRF-Token") |
| if not csrf_token or not self._validate_csrf_token(csrf_token): |
| security_logger.warning( |
| f"CSRF token validation failed for: {request.method} {request.url.path}" |
| ) |
| return JSONResponse( |
| status_code=403, |
| content={ |
| "error": { |
| "type": "csrf_token_invalid", |
| "message": "Invalid or missing CSRF token" |
| } |
| } |
| ) |
|
|
| return await call_next(request) |
|
|
| def generate_csrf_token(self, session_id: str) -> str: |
| """Generate CSRF token for session""" |
| token = secrets.token_urlsafe(32) |
| expiry = time.time() + self.token_expiry |
|
|
| self.csrf_tokens[token] = { |
| "session_id": session_id, |
| "expiry": expiry |
| } |
|
|
| return token |
|
|
| def _validate_csrf_token(self, token: str) -> bool: |
| """Validate CSRF token""" |
| if token not in self.csrf_tokens: |
| return False |
|
|
| token_data = self.csrf_tokens[token] |
|
|
| |
| if time.time() > token_data["expiry"]: |
| del self.csrf_tokens[token] |
| return False |
|
|
| return True |
|
|
|
|
| def setup_security_middleware(app): |
| """Setup all security middleware""" |
| |
| app.add_middleware(SecurityHeadersMiddleware) |
| app.add_middleware(CSRFProtectionMiddleware) |
| app.add_middleware(InputValidationMiddleware) |
| app.add_middleware(RateLimitMiddleware, requests_per_minute=120, burst_size=20) |
|
|
|
|
| |
| def hash_password(password: str) -> str: |
| """Hash password using secure bcrypt implementation""" |
| return secure_hash_password(password) |
|
|
|
|
| def generate_api_key() -> str: |
| """Generate secure API key""" |
| return secrets.token_urlsafe(32) |
|
|
|
|
| def validate_email(email: str) -> bool: |
| """Validate email format""" |
| pattern = r'^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$' |
| return re.match(pattern, email) is not None |
|
|
|
|
| def sanitize_input(input_str: str) -> str: |
| """Sanitize user input""" |
| |
| clean = re.sub(r'<[^>]+>', '', input_str) |
| |
| clean = re.sub(r'[<>"\']', '', clean) |
| return clean.strip() |