|
|
| """Tests for per-user rate limiting."""
|
|
|
| import os
|
| import sys
|
| import tempfile
|
| import unittest
|
| from unittest.mock import MagicMock
|
|
|
| sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
|
|
|
| _test_db = tempfile.NamedTemporaryFile(delete=False, suffix=".db")
|
| _test_db.close()
|
|
|
| os.environ["DB_BACKEND"] = "sqlite"
|
| os.environ["SQLITE_DB_PATH"] = _test_db.name
|
| os.environ["RATE_LIMIT_ENABLED"] = "true"
|
| os.environ["RATE_LIMIT_PER_HOUR"] = "3"
|
| os.environ["RATE_LIMIT_PER_DAY"] = "5"
|
| os.environ["RATE_LIMIT_BATCH_PER_HOUR"] = "2"
|
| os.environ["DEMO_ACCESS_PASSWORD"] = "test-access-password"
|
| os.environ["DEMO_PASSWORD_REQUEST_URL"] = "https://bartoszlenart.com"
|
|
|
|
|
| sys.modules.setdefault("psycopg2", MagicMock())
|
| sys.modules.setdefault("psycopg2.extras", MagicMock())
|
|
|
| from database import DatabaseManager
|
| from rate_limiter import RateLimiter, get_user_id_from_request
|
|
|
|
|
| class TestRateLimiter(unittest.TestCase):
|
| def setUp(self):
|
| self.db = DatabaseManager()
|
| self.limiter = RateLimiter()
|
|
|
| def test_allows_requests_under_limit(self):
|
| result = self.limiter.check_and_record("user-a", "process")
|
| self.assertTrue(result.allowed)
|
| self.assertEqual(result.remaining_hour, 2)
|
|
|
| def test_blocks_after_hourly_limit(self):
|
| for _ in range(3):
|
| self.limiter.check_and_record("user-b", "process")
|
|
|
| blocked = self.limiter.check("user-b", "process")
|
| self.assertFalse(blocked.allowed)
|
| self.assertIn("Rate limit reached", blocked.message)
|
| self.assertIn("https://bartoszlenart.com", blocked.message)
|
|
|
| def test_password_allows_request_after_hourly_limit(self):
|
| for _ in range(3):
|
| self.limiter.check_and_record("user-password", "process")
|
|
|
| allowed = self.limiter.check(
|
| "user-password", "process", access_password="test-access-password"
|
| )
|
| self.assertTrue(allowed.allowed)
|
|
|
| def test_wrong_password_does_not_bypass_limit(self):
|
| for _ in range(3):
|
| self.limiter.check_and_record("user-wrong-password", "process")
|
|
|
| blocked = self.limiter.check(
|
| "user-wrong-password", "process", access_password="wrong"
|
| )
|
| self.assertFalse(blocked.allowed)
|
|
|
| def test_batch_limit_is_separate(self):
|
| for _ in range(2):
|
| self.limiter.check_and_record("user-c", "batch")
|
|
|
| blocked = self.limiter.check("user-c", "batch")
|
| self.assertFalse(blocked.allowed)
|
| self.assertIn("Batch demo limit", blocked.message)
|
|
|
| allowed = self.limiter.check("user-c", "process")
|
| self.assertTrue(allowed.allowed)
|
|
|
| def test_disabled_when_env_false(self):
|
| os.environ["RATE_LIMIT_ENABLED"] = "false"
|
| disabled_limiter = RateLimiter()
|
| result = disabled_limiter.check("user-d", "process")
|
| self.assertTrue(result.allowed)
|
| os.environ["RATE_LIMIT_ENABLED"] = "true"
|
|
|
| def test_get_user_id_from_forwarded_header(self):
|
| request = MagicMock()
|
| request.headers = {"x-forwarded-for": "203.0.113.10, 70.41.3.18"}
|
| request.client = MagicMock(host="127.0.0.1")
|
| self.assertEqual(get_user_id_from_request(request), "203.0.113.10")
|
|
|
| def test_get_user_id_from_client_host(self):
|
| request = MagicMock()
|
| request.headers = {}
|
| request.client = MagicMock(host="192.168.1.5")
|
| self.assertEqual(get_user_id_from_request(request), "192.168.1.5")
|
|
|
|
|
| if __name__ == "__main__":
|
| unittest.main()
|
|
|