fhirflame / tests /test_rate_limiter.py
grasant's picture
Add password gate after demo usage limits
6f9bb4a verified
Raw
History Blame Contribute Delete
3.68 kB
#!/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()