"""Tests for rate limiting on contact and rental endpoints (T05).""" import pytest from unittest.mock import AsyncMock, patch, MagicMock from httpx import AsyncClient, ASGITransport from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from app.database import Base, get_db from app.models import EquipmentCache VALID_CONTACT = { "name": "Rate Limit", "email": "rate@example.com", "phone": "+49 123", "message": "Rate limit test", "privacy_consent": True, } VALID_RENTAL = { "event_name": "Rate Limit Fest", "date_start": "2026-09-01", "date_end": "2026-09-02", "location": "Hamburg", "contact_name": "RL User", "contact_email": "rl@example.com", "items": [{"equipment_id": 1, "quantity": 1}], } def _make_rate_limit_client(test_engine, rate_counts: list[int]): """Create a client whose cache.incr_rate returns sequential values from rate_counts.""" session_maker = async_sessionmaker(test_engine, class_=AsyncSession, expire_on_commit=False) call_index = {"i": 0} async def override_get_db(): async with session_maker() as session: yield session async def mock_incr_rate(key: str, window: int = 60) -> int: idx = call_index["i"] if idx < len(rate_counts): val = rate_counts[idx] else: val = rate_counts[-1] + 1 call_index["i"] += 1 return val mock_cache = MagicMock() mock_cache.get = AsyncMock(return_value=None) mock_cache.set = AsyncMock(return_value=None) mock_cache.delete_pattern = AsyncMock(return_value=0) mock_cache.incr_rate = mock_incr_rate mock_cache.connect = AsyncMock(return_value=None) mock_cache._redis = MagicMock() mock_cache._redis.ping = AsyncMock(return_value=True) async def _yield_client(): with patch("app.cache.cache", mock_cache), \ patch("app.routers.equipment.cache", mock_cache), \ patch("app.routers.admin.cache", mock_cache), \ patch("app.routers.rental_requests.cache", mock_cache), \ patch("app.routers.contact.cache", mock_cache), \ patch("app.services.sync_service.cache", mock_cache): from app.main import app app.dependency_overrides[get_db] = override_get_db transport = ASGITransport(app=app) async with AsyncClient(transport=transport, base_url="http://test") as ac: yield ac app.dependency_overrides.clear() return _yield_client @pytest.mark.asyncio async def test_contact_rate_limit_allows_5_blocks_6th(test_engine): """5 requests OK, 6th request returns 429 for contact endpoint.""" rate_counts = [1, 2, 3, 4, 5, 6] client_gen = _make_rate_limit_client(test_engine, rate_counts) async for client in client_gen(): with patch( "app.services.email_service.EmailService.send_contact_email", new_callable=AsyncMock, return_value=True, ): for i in range(5): resp = await client.post("/api/contact", json=VALID_CONTACT) assert resp.status_code == 200, f"Request {i+1} should succeed" resp = await client.post("/api/contact", json=VALID_CONTACT) assert resp.status_code == 429 @pytest.mark.asyncio async def test_rental_rate_limit_allows_5_blocks_6th(test_engine, seeded_equipment): """5 requests OK, 6th request returns 429 for rental-requests endpoint.""" rate_counts = [1, 2, 3, 4, 5, 6] client_gen = _make_rate_limit_client(test_engine, rate_counts) async for client in client_gen(): with patch( "app.services.rentman_service.RentmanService.create_project_request", new_callable=AsyncMock, return_value={"id": "rl-test"}, ), patch( "app.services.rentman_service.RentmanService.add_equipment_to_request", new_callable=AsyncMock, return_value={"id": "eq-rl"}, ), patch( "app.services.email_service.EmailService.send_rental_confirmation", new_callable=AsyncMock, return_value=True, ): for i in range(5): resp = await client.post("/api/rental-requests", json=VALID_RENTAL) assert resp.status_code == 201, f"Request {i+1} should succeed" resp = await client.post("/api/rental-requests", json=VALID_RENTAL) assert resp.status_code == 429 @pytest.mark.asyncio async def test_contact_rate_window_reset(test_engine): """Rate window resets after TTL: first 5 OK, then 429, then after reset 5 more OK.""" rate_counts = [1, 2, 3, 4, 5, 6, 1, 2] client_gen = _make_rate_limit_client(test_engine, rate_counts) async for client in client_gen(): with patch( "app.services.email_service.EmailService.send_contact_email", new_callable=AsyncMock, return_value=True, ): for i in range(5): resp = await client.post("/api/contact", json=VALID_CONTACT) assert resp.status_code == 200 resp = await client.post("/api/contact", json=VALID_CONTACT) assert resp.status_code == 429 # After window reset, counter starts at 1 again resp = await client.post("/api/contact", json=VALID_CONTACT) assert resp.status_code == 200 resp = await client.post("/api/contact", json=VALID_CONTACT) assert resp.status_code == 200