from __future__ import annotations from datetime import datetime, timezone from unittest.mock import MagicMock import pytest from fastapi import HTTPException from sqlalchemy import create_engine, select from sqlalchemy.orm import Session from app.database import Base from app.main import _check_rate_limit from app.models import SubmissionAttempt def test_rate_limit_is_persistent_and_does_not_store_raw_client() -> None: engine = create_engine("sqlite://") Base.metadata.create_all(engine) with Session(engine) as db: for _ in range(5): mock_request = MagicMock() mock_request.client.host = "203.0.113.42" mock_request.headers.get.return_value = None _check_rate_limit(mock_request, db) with pytest.raises(HTTPException) as blocked: mock_request = MagicMock() mock_request.client.host = "203.0.113.42" mock_request.headers.get.return_value = None _check_rate_limit(mock_request, db) assert blocked.value.status_code == 429 attempts = list(db.scalars(select(SubmissionAttempt))) assert len(attempts) == 5 assert all(item.client_hash != "203.0.113.42" and len(item.client_hash) == 64 for item in attempts) assert all(item.created_at.replace(tzinfo=timezone.utc) <= datetime.now(timezone.utc) for item in attempts) def test_rate_limit_uses_forwarded_for_header() -> None: engine = create_engine("sqlite://") Base.metadata.create_all(engine) with Session(engine) as db: mock_real = MagicMock() mock_real.client.host = "10.0.0.1" mock_real.headers.get.return_value = "198.51.100.10" for _ in range(5): _check_rate_limit(mock_real, db) with pytest.raises(HTTPException) as blocked: mock_other = MagicMock() mock_other.client.host = "10.0.0.2" mock_other.headers.get.return_value = "198.51.100.10" _check_rate_limit(mock_other, db) assert blocked.value.status_code == 429