#!/usr/bin/env python3 """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" # Allow sqlite-only test runs without PostgreSQL drivers installed. 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()