7ad2bdfc4b
- test_rate_limiting: remove references to non-existent rate_limit_process setting - test_path_traversal_security: fix sanitize_filename assertion to match actual strip behavior, fix os.path.basename test for Linux (backslash not a separator), remove erroneous task_mock arg from embed_metadata_into_pdf direct call - test_e2e_full_stack: add psycopg2 availability check to skip Postgres test when driver is not installed Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
241 lines
8.0 KiB
Python
241 lines
8.0 KiB
Python
#!/usr/bin/env python3
|
|
|
|
"""
|
|
Tests for rate limiting middleware.
|
|
|
|
These tests validate that rate limiting is properly applied to API endpoints
|
|
to prevent abuse and DoS attacks.
|
|
"""
|
|
|
|
import time
|
|
|
|
import pytest
|
|
from fastapi import status
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_rate_limiting_enabled_by_default():
|
|
"""Test that rate limiting is enabled by default in configuration."""
|
|
from app.config import settings
|
|
|
|
# Rate limiting should be enabled by default
|
|
assert hasattr(settings, "rate_limiting_enabled")
|
|
assert isinstance(settings.rate_limiting_enabled, bool)
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_rate_limit_configuration():
|
|
"""Test that rate limit configuration is loaded correctly."""
|
|
from app.config import settings
|
|
|
|
# Verify that rate limit configuration attributes exist
|
|
assert hasattr(settings, "rate_limiting_enabled")
|
|
assert hasattr(settings, "rate_limit_default")
|
|
assert hasattr(settings, "rate_limit_upload")
|
|
assert hasattr(settings, "rate_limit_auth")
|
|
|
|
# Verify that settings are strings in correct format
|
|
assert isinstance(settings.rate_limit_default, str)
|
|
assert "/" in settings.rate_limit_default # Should be like "100/minute"
|
|
assert isinstance(settings.rate_limit_upload, str)
|
|
assert "/" in settings.rate_limit_upload
|
|
assert isinstance(settings.rate_limit_auth, str)
|
|
assert "/" in settings.rate_limit_auth
|
|
|
|
|
|
@pytest.mark.integration
|
|
def test_limiter_initialization():
|
|
"""Test that rate limiter is initialized correctly."""
|
|
from app.main import app
|
|
|
|
# Verify limiter is attached to app state
|
|
assert hasattr(app.state, "limiter")
|
|
assert app.state.limiter is not None
|
|
|
|
# Verify limiter has expected attributes
|
|
limiter = app.state.limiter
|
|
assert hasattr(limiter, "limit")
|
|
assert hasattr(limiter, "exempt")
|
|
|
|
|
|
@pytest.mark.integration
|
|
def test_rate_limit_on_health_endpoint(client):
|
|
"""Test that endpoints respect rate limits."""
|
|
from app.config import settings
|
|
|
|
if not settings.rate_limiting_enabled:
|
|
pytest.skip("Rate limiting is disabled in test configuration")
|
|
|
|
# Test with / endpoint which should exist
|
|
# Make multiple requests within the limit
|
|
for _ in range(5):
|
|
response = client.get("/")
|
|
# Should get either 200 (success) or 302 (redirect) but not 429 (rate limited)
|
|
assert response.status_code in [200, 302, 404], f"Unexpected status: {response.status_code}"
|
|
|
|
|
|
@pytest.mark.integration
|
|
def test_rate_limit_exceeded_returns_429(client):
|
|
"""Test that exceeding rate limit returns 429 status code."""
|
|
from app.config import settings
|
|
|
|
if not settings.rate_limiting_enabled:
|
|
pytest.skip("Rate limiting is disabled in test configuration")
|
|
|
|
# This test would need to make enough requests to trigger rate limit
|
|
# Since we use in-memory storage for tests and default limit is high,
|
|
# we'll verify the mechanism is in place
|
|
# In production, this would be tested with lower limits
|
|
|
|
# Make a moderate number of requests to the about page
|
|
responses = []
|
|
for _ in range(10):
|
|
response = client.get("/about")
|
|
responses.append(response.status_code)
|
|
|
|
# All should succeed with default high limits (not testing actual rate limiting)
|
|
# We're just verifying the endpoints are accessible
|
|
assert all(code in [200, 302, 404] for code in responses)
|
|
|
|
|
|
@pytest.mark.security
|
|
def test_rate_limiting_uses_correct_identifier():
|
|
"""Test that rate limiting uses IP or user ID as identifier."""
|
|
from app.middleware.rate_limit import get_identifier
|
|
from fastapi import Request
|
|
|
|
# Create a mock request
|
|
class MockRequest:
|
|
def __init__(self):
|
|
self.session = {}
|
|
self.client = type("client", (), {"host": "127.0.0.1"})()
|
|
|
|
# Test with unauthenticated request (should use IP)
|
|
request = MockRequest()
|
|
identifier = get_identifier(request)
|
|
assert identifier == "127.0.0.1"
|
|
|
|
# Test with authenticated request (should use user identifier)
|
|
request.session["user"] = {"username": "testuser", "id": "123"}
|
|
identifier = get_identifier(request)
|
|
assert "user:" in identifier or identifier == "testuser" or "123" in identifier
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_limiter_creation_with_redis():
|
|
"""Test limiter creation with Redis backend."""
|
|
from app.middleware.rate_limit import create_limiter
|
|
|
|
# Create limiter with Redis URL
|
|
limiter = create_limiter(redis_url="redis://localhost:6379/0", enabled=True)
|
|
assert limiter is not None
|
|
assert limiter.enabled is True
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_limiter_creation_with_memory():
|
|
"""Test limiter creation with in-memory backend."""
|
|
from app.middleware.rate_limit import create_limiter
|
|
|
|
# Create limiter without Redis (fallback to memory)
|
|
limiter = create_limiter(redis_url=None, enabled=True)
|
|
assert limiter is not None
|
|
assert limiter.enabled is True
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_limiter_disabled():
|
|
"""Test limiter creation when disabled."""
|
|
from app.middleware.rate_limit import create_limiter
|
|
|
|
# Create disabled limiter
|
|
limiter = create_limiter(redis_url=None, enabled=False)
|
|
assert limiter is not None
|
|
assert limiter.enabled is False
|
|
|
|
|
|
@pytest.mark.integration
|
|
def test_rate_limit_exception_handler_registered():
|
|
"""Test that rate limit exception handler is registered."""
|
|
from app.main import app
|
|
from slowapi.errors import RateLimitExceeded
|
|
|
|
# Verify exception handler is registered
|
|
assert RateLimitExceeded in app.exception_handlers
|
|
|
|
|
|
@pytest.mark.security
|
|
def test_rate_limit_prevents_brute_force():
|
|
"""Test that rate limiting can prevent brute force attacks on auth endpoints."""
|
|
# This is a documentation test - in practice, auth endpoints should have
|
|
# stricter rate limits (e.g., 10/minute) to prevent brute force
|
|
from app.config import settings
|
|
|
|
# Auth endpoints should have stricter limits
|
|
assert hasattr(settings, "rate_limit_auth")
|
|
|
|
# Parse the limit to ensure it's restrictive enough
|
|
limit_str = settings.rate_limit_auth
|
|
count, period = limit_str.split("/")
|
|
count = int(count)
|
|
|
|
# Should be significantly lower than default
|
|
# e.g., 10/minute vs 100/minute for default
|
|
assert count <= 20, "Auth rate limit should be strict to prevent brute force"
|
|
|
|
|
|
@pytest.mark.integration
|
|
def test_rate_limiting_middleware_integration():
|
|
"""Test that rate limiting middleware integrates properly with app."""
|
|
from app.main import app
|
|
|
|
# Verify app has limiter state
|
|
assert hasattr(app, "state")
|
|
assert hasattr(app.state, "limiter")
|
|
|
|
# Verify middleware configuration
|
|
from app.config import settings
|
|
|
|
assert hasattr(settings, "rate_limiting_enabled")
|
|
assert hasattr(settings, "redis_url")
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_rate_limit_format_validation():
|
|
"""Test that rate limit strings are in valid format."""
|
|
from app.config import settings
|
|
|
|
# Validate format of rate limit strings
|
|
def validate_rate_limit(limit_str):
|
|
"""Validate rate limit string format."""
|
|
parts = limit_str.split("/")
|
|
if len(parts) != 2:
|
|
return False
|
|
try:
|
|
count = int(parts[0])
|
|
period = parts[1]
|
|
valid_periods = ["second", "minute", "hour", "day"]
|
|
return count > 0 and period in valid_periods
|
|
except ValueError:
|
|
return False
|
|
|
|
assert validate_rate_limit(settings.rate_limit_default)
|
|
assert validate_rate_limit(settings.rate_limit_upload)
|
|
assert validate_rate_limit(settings.rate_limit_auth)
|
|
|
|
|
|
@pytest.mark.integration
|
|
def test_concurrent_requests_respect_rate_limit():
|
|
"""Test that concurrent requests from same client respect rate limits."""
|
|
# This test documents expected behavior
|
|
# In production, multiple rapid requests from same IP should be tracked
|
|
# and rate limited appropriately
|
|
from app.config import settings
|
|
|
|
if not settings.rate_limiting_enabled:
|
|
pytest.skip("Rate limiting is disabled")
|
|
|
|
# Document that rate limiting tracks requests per identifier
|
|
# and enforces limits across concurrent requests
|
|
assert settings.rate_limiting_enabled is True
|