File size: 3,675 Bytes
6f9bb4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
#!/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()