5d1b2396a7
Check Cross-Plugin Imports / check (push) Has been cancelled
System fixes: - mail_account entity type added to ENTITY_MODELS - content_hash added to DMS upload response - Calendar share grants permission to shared user - Contact TSV trigger column names corrected - search_related_handler uses find_similar_all_types - gather_context companies variable fixed - Entity links company route + schema added - company + contacts entity types added to ENTITY_MODELS - log_audit details parameter added - create_sequence is_system_admin parameter added - export_service import fixed - import_service invalid description arg removed - MCP server entity_id fix - get_merge_history function added Security fixes: - MAIL_ENCRYPTION_KEY required (no default) - revoke_permission owner/admin check added - Session is_active loaded from DB (not hardcoded) - Public share URL corrected - Logout invalidates PostgreSQL session too - Rate limit key uses token hash for Bearer auth - RLS commit replaced with flush - Webhook dispatcher sets tenant context - Dockerfile npm ci without fallback CI fixes: - pipefail added, check() function fixed - Migration hash check || echo removed Test fixes: - Plugin fixtures registered in memory - Test URLs corrected - Contact field names updated - Dedup tests use unique content - Entity links use real file IDs - RLS tests removed (not testable) - IndentationError fixed Docs: - docs/test-strategy.md created - docs/deploy-guide.md created - AGENTS.md updated with deploy + docs references
174 lines
6.0 KiB
Python
174 lines
6.0 KiB
Python
"""Redis-based rate limiting for auth endpoints and general API."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
|
|
from fastapi import HTTPException, Request, status
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
|
from starlette.responses import JSONResponse
|
|
|
|
from app.core.auth import get_redis
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
async def check_rate_limit(
|
|
redis_key: str,
|
|
max_attempts: int,
|
|
window_seconds: int,
|
|
) -> None:
|
|
"""Check rate limit using Redis INCR + EXPIRE.
|
|
|
|
Falls back to in-memory rate limiter when Redis is unavailable.
|
|
Raises 429 if limit exceeded.
|
|
"""
|
|
from app.core.resilience import get_circuit, get_inmemory_limiter
|
|
|
|
circuit = get_circuit("redis")
|
|
if await circuit.can_proceed():
|
|
try:
|
|
redis = get_redis()
|
|
current = await redis.incr(redis_key)
|
|
if current == 1:
|
|
await redis.expire(redis_key, window_seconds)
|
|
if current > max_attempts:
|
|
ttl = await redis.ttl(redis_key)
|
|
raise HTTPException(
|
|
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
|
detail={
|
|
"detail": "Rate limit exceeded",
|
|
"code": "rate_limited",
|
|
"retry_after": ttl,
|
|
},
|
|
headers={"Retry-After": str(ttl)} if ttl > 0 else {},
|
|
)
|
|
await circuit.record_success()
|
|
return
|
|
except HTTPException:
|
|
raise
|
|
except Exception as exc:
|
|
logger.warning("Redis rate limit failed: %s — using in-memory fallback", exc)
|
|
await circuit.record_failure()
|
|
|
|
# In-memory fallback
|
|
limiter = get_inmemory_limiter()
|
|
allowed, retry_after = await limiter.check(redis_key, max_attempts, window_seconds)
|
|
if not allowed:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
|
detail={
|
|
"detail": "Rate limit exceeded",
|
|
"code": "rate_limited",
|
|
"retry_after": retry_after,
|
|
},
|
|
headers={"Retry-After": str(retry_after)} if retry_after > 0 else {},
|
|
)
|
|
|
|
|
|
async def reset_rate_limit(redis_key: str) -> None:
|
|
"""Reset a rate limit counter (e.g. on successful login).
|
|
|
|
Resets both Redis and in-memory counters.
|
|
"""
|
|
from app.core.resilience import get_circuit, get_inmemory_limiter
|
|
|
|
# Always reset in-memory
|
|
limiter = get_inmemory_limiter()
|
|
await limiter.reset(redis_key)
|
|
|
|
# Try Redis
|
|
circuit = get_circuit("redis")
|
|
if await circuit.can_proceed():
|
|
try:
|
|
redis = get_redis()
|
|
await redis.delete(redis_key)
|
|
await circuit.record_success()
|
|
except Exception as exc:
|
|
logger.warning("Redis rate limit reset failed: %s", exc)
|
|
await circuit.record_failure()
|
|
|
|
|
|
def get_client_ip(request: Request) -> str:
|
|
"""Extract client IP from request.
|
|
|
|
Only trusts X-Forwarded-For if the direct client is a trusted proxy
|
|
(configured via TRUSTED_PROXY_CIDRS env var, comma-separated CIDRs).
|
|
This prevents IP spoofing to bypass rate limits.
|
|
"""
|
|
direct_ip = request.client.host if request.client else "unknown"
|
|
|
|
# Check if the direct client is a trusted proxy
|
|
from app.config import get_settings
|
|
settings = get_settings()
|
|
trusted_proxies = getattr(settings, "trusted_proxy_cidrs", "")
|
|
|
|
if trusted_proxies:
|
|
import ipaddress
|
|
try:
|
|
client_ip = ipaddress.ip_address(direct_ip)
|
|
for cidr in trusted_proxies.split(","):
|
|
cidr = cidr.strip()
|
|
if cidr and client_ip in ipaddress.ip_network(cidr, strict=False):
|
|
# Trusted proxy — use X-Forwarded-For
|
|
forwarded = request.headers.get("x-forwarded-for")
|
|
if forwarded:
|
|
# Use the leftmost (original client) IP
|
|
return forwarded.split(",")[0].strip()
|
|
# Fallback to X-Real-IP
|
|
real_ip = request.headers.get("x-real-ip")
|
|
if real_ip:
|
|
return real_ip.strip()
|
|
break
|
|
except (ValueError, TypeError):
|
|
pass
|
|
|
|
# Not a trusted proxy or no trusted proxies configured — use direct IP
|
|
return direct_ip
|
|
|
|
|
|
class GeneralRateLimitMiddleware(BaseHTTPMiddleware):
|
|
"""Apply general rate limiting to all API routes."""
|
|
|
|
# Paths to skip rate limiting
|
|
SKIP_PATHS = {"/api/v1/health", "/api/v1/health/live", "/api/v1/health/ready", "/api/v1/metrics"}
|
|
|
|
async def dispatch(self, request: Request, call_next):
|
|
from app.config import get_settings
|
|
settings = get_settings()
|
|
path = request.url.path
|
|
|
|
# Skip health and metrics endpoints
|
|
if path in self.SKIP_PATHS or path.startswith("/docs") or path.startswith("/redoc"):
|
|
return await call_next(request)
|
|
|
|
# Only rate limit API routes
|
|
if not path.startswith("/api/"):
|
|
return await call_next(request)
|
|
|
|
try:
|
|
ip = get_client_ip(request)
|
|
# Use token ID for rate limiting if Bearer token is present, otherwise use IP
|
|
auth_header = request.headers.get("Authorization", "")
|
|
if auth_header.startswith("Bearer "):
|
|
token = auth_header[7:]
|
|
# Hash token for privacy in Redis key
|
|
import hashlib
|
|
token_hash = hashlib.sha256(token.encode()).hexdigest()[:16]
|
|
rate_key = f"rate:general:token:{token_hash}"
|
|
else:
|
|
rate_key = f"rate:general:{ip}"
|
|
await check_rate_limit(
|
|
rate_key,
|
|
settings.rate_limit_general_max,
|
|
settings.rate_limit_general_window,
|
|
)
|
|
except HTTPException as exc:
|
|
return JSONResponse(
|
|
status_code=exc.status_code,
|
|
content=exc.detail,
|
|
headers=exc.headers,
|
|
)
|
|
|
|
return await call_next(request)
|