52 lines
2.0 KiB
Python
52 lines
2.0 KiB
Python
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
|