Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
76 changes: 72 additions & 4 deletions dashboard/backend/middleware/rate_limit.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,18 +8,28 @@
- Redis-backed for distributed rate limiting
"""

import os
import time
from typing import Optional, Callable
from fastapi import Request, Response, HTTPException
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.responses import JSONResponse
import structlog
from prometheus_client import Counter

from core.rate_limit_config import get_rate_limit, get_burst_limit, get_endpoint_limit
from services.cache_service import get_redis_client

logger = structlog.get_logger()

# Backing-store (redis) failures inside the rate limiter are never silent:
# every failure increments this counter (exposed via /metrics, alertable).
RATE_LIMIT_BACKEND_FAILURES = Counter(
"rate_limit_backend_failures",
"Rate limit checks that failed on the backing store (e.g. redis errors)",
["fail_mode"],
)


class RateLimiter:
"""
Expand All @@ -32,15 +42,25 @@ class RateLimiter:
- Sliding window for accurate rate limiting
"""

def __init__(self, redis_client=None):
def __init__(self, redis_client=None, fail_mode: Optional[str] = None):
"""
Initialize rate limiter

Args:
redis_client: Redis client (optional, will use default if not provided)
fail_mode: Behavior on backing-store failure ("open"|"closed").
Defaults to env RATE_LIMIT_FAIL_MODE, else "open".
Valid fail modes (env RATE_LIMIT_FAIL_MODE):
- "open" (default): allow requests if the backing store fails, but
mark the response as degraded (header + metric + log event).
- "closed": deny requests with 503 while the backing store is down.
"""
self.redis = redis_client or get_redis_client()
self.prefix = "ratelimit:"
self.fail_mode = (fail_mode or os.getenv("RATE_LIMIT_FAIL_MODE", "open")).strip().lower()
if self.fail_mode not in ("open", "closed"):
logger.warning("Invalid RATE_LIMIT_FAIL_MODE, falling back to open", value=self.fail_mode)
self.fail_mode = "open"

async def is_allowed(
self,
Expand Down Expand Up @@ -74,6 +94,10 @@ async def is_allowed(
endpoint_limit,
window=60 # 1 minute
)
if info.get("degraded"):
# Backing store down: short-circuit so each request produces
# exactly one failure event, not one per limit bucket.
return is_allowed, info
if not is_allowed:
return False, info

Expand All @@ -84,6 +108,8 @@ async def is_allowed(
burst_limit,
window=60 # 1 minute
)
if burst_info.get("degraded"):
return is_allowed, burst_info
if not is_allowed:
return False, burst_info

Expand All @@ -94,6 +120,8 @@ async def is_allowed(
hourly_limit,
window=3600 # 1 hour
)
if hourly_info.get("degraded"):
return is_allowed, hourly_info
if not is_allowed:
return False, hourly_info

Expand Down Expand Up @@ -150,9 +178,30 @@ async def _check_limit(
return is_allowed, info

except Exception as e:
logger.error("Rate limit check failed", error=str(e), key=key)
# Fail open - allow request if Redis fails
return True, {"limit": limit, "remaining": limit, "error": str(e)}
# Backing-store failure — never silent: distinct log event + metric.
# Details stay server-side (no error strings in responses).
try:
RATE_LIMIT_BACKEND_FAILURES.labels(fail_mode=self.fail_mode).inc()
except Exception:
pass
logger.error(
"rate_limit_backend_failure",
fail_mode=self.fail_mode,
error_type=type(e).__name__,
error=str(e),
key=key,
)
info = {
"limit": limit,
"remaining": limit,
"reset": int(current_time + window),
"window": window,
"degraded": True,
}
if self.fail_mode == "closed":
return False, info
# Fail open (default): allow request but flag degradation
return True, info


class RateLimitMiddleware(BaseHTTPMiddleware):
Expand Down Expand Up @@ -197,14 +246,28 @@ async def dispatch(self, request: Request, call_next: Callable) -> Response:
endpoint=request.url.path
)

degraded = bool(rate_info.get("degraded"))

# Add rate limit headers
headers = {
"X-RateLimit-Limit": str(rate_info.get("limit", 0)),
"X-RateLimit-Remaining": str(rate_info.get("remaining", 0)),
"X-RateLimit-Reset": str(rate_info.get("reset", 0))
}
if degraded:
# Externally visible signal that rate limiting is currently
# unavailable (backing store failure).
headers["X-RateLimit-Mode"] = "degraded"

if not is_allowed:
if degraded:
# RATE_LIMIT_FAIL_MODE=closed: backing store down. This is a
# service degradation, NOT a client abuse signal — 503, not 429.
return JSONResponse(
status_code=503,
content={"detail": "Service temporarily degraded, please retry later"},
headers={**headers, "Retry-After": "30"}
)
# Rate limit exceeded
logger.warning(
"Rate limit exceeded",
Expand Down Expand Up @@ -276,6 +339,11 @@ async def endpoint(request: Request, _: None = Depends(check_rate_limit)):
)

if not is_allowed:
if rate_info.get("degraded"):
raise HTTPException(
status_code=503,
detail="Service temporarily degraded, please retry later"
)
raise HTTPException(
status_code=429,
detail={
Expand Down
99 changes: 97 additions & 2 deletions dashboard/backend/tests/middleware/test_rate_limit.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,7 +180,7 @@ async def test_is_allowed_burst_limit(self, rate_limiter, mock_redis):

@pytest.mark.asyncio
async def test_redis_failure_fails_open(self, rate_limiter, mock_redis):
"""Should allow request if Redis fails"""
"""Should allow request if Redis fails (default open mode), flagged degraded"""
mock_redis.zcard = AsyncMock(side_effect=Exception("Redis error"))

is_allowed, info = await rate_limiter.is_allowed(
Expand All @@ -190,7 +190,102 @@ async def test_redis_failure_fails_open(self, rate_limiter, mock_redis):
)

assert is_allowed is True
assert "error" in info
assert info.get("degraded") is True
# Error details must stay server-side (S-3 discipline: no leakage)
assert "error" not in info


class TestFailModeBehavior:
"""Behavior on backing-store failure (card 4f9fb443, audit OBS-1)"""

@pytest.fixture
def failing_redis(self):
redis = AsyncMock()
redis.zremrangebyscore = AsyncMock(side_effect=Exception("Event loop is closed"))
return redis

@pytest.mark.asyncio
async def test_fail_open_marks_degraded(self, failing_redis):
"""Open mode: request allowed, degraded flag set"""
limiter = RateLimiter(redis_client=failing_redis, fail_mode="open")
is_allowed, info = await limiter.is_allowed("user:1", "free", "/api/v1/test")
assert is_allowed is True
assert info.get("degraded") is True

@pytest.mark.asyncio
async def test_fail_closed_denies(self, failing_redis):
"""Closed mode: request denied, degraded flag set"""
limiter = RateLimiter(redis_client=failing_redis, fail_mode="closed")
is_allowed, info = await limiter.is_allowed("user:1", "free", "/api/v1/test")
assert is_allowed is False
assert info.get("degraded") is True

@pytest.mark.asyncio
async def test_backend_failure_short_circuits_single_event(self, failing_redis):
"""One backing-store failure per request: first redis error short-circuits"""
limiter = RateLimiter(redis_client=failing_redis, fail_mode="open")
await limiter.is_allowed("user:1", "pro", "/api/v1/auth/login")
# Endpoint limit bucket is the first check → exactly 1 redis call, not 3
assert failing_redis.zremrangebyscore.await_count == 1

@pytest.mark.asyncio
async def test_failure_metric_incremented(self, failing_redis):
"""Every backing-store failure increments the Prometheus counter"""
from prometheus_client import REGISTRY

def sample(mode):
return REGISTRY.get_sample_value(
"rate_limit_backend_failures_total", {"fail_mode": mode}
) or 0

before_open = sample("open")
before_closed = sample("closed")

await RateLimiter(redis_client=failing_redis, fail_mode="open").is_allowed(
"user:1", "free", "/api/v1/test")
await RateLimiter(redis_client=failing_redis, fail_mode="closed").is_allowed(
"user:1", "free", "/api/v1/test")

assert sample("open") == before_open + 1
assert sample("closed") == before_closed + 1

def test_invalid_fail_mode_falls_back_open(self, failing_redis):
"""Unknown RATE_LIMIT_FAIL_MODE value falls back to open"""
limiter = RateLimiter(redis_client=failing_redis, fail_mode="yolo")
assert limiter.fail_mode == "open"

def test_middleware_open_mode_200_with_degraded_header(self, failing_redis):
"""Middleware: open mode keeps serving but exposes X-RateLimit-Mode: degraded"""
app = FastAPI()
app.add_middleware(
RateLimitMiddleware,
rate_limiter=RateLimiter(redis_client=failing_redis, fail_mode="open"),
)

@app.get("/test")
async def test_endpoint():
return {"status": "ok"}

response = TestClient(app).get("/test")
assert response.status_code == 200
assert response.headers.get("X-RateLimit-Mode") == "degraded"

def test_middleware_closed_mode_503_no_leak(self, failing_redis):
"""Middleware: closed mode returns 503 + Retry-After, no backend error leaked"""
app = FastAPI()
app.add_middleware(
RateLimitMiddleware,
rate_limiter=RateLimiter(redis_client=failing_redis, fail_mode="closed"),
)

@app.get("/test")
async def test_endpoint():
return {"status": "ok"}

response = TestClient(app).get("/test")
assert response.status_code == 503
assert response.headers.get("Retry-After") == "30"
assert "Event loop is closed" not in response.text


class TestRateLimitMiddleware:
Expand Down
Loading