File size: 3,630 Bytes
99a7ebb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
103
104
import unittest
from unittest import mock

import requests

from services.register.mail_provider import CloudflareTempMailProvider, YydsMailProvider


class FakeResponse:
    def __init__(self, status_code=200, data=None, text=""):
        self.status_code = status_code
        self._data = data if data is not None else {}
        self.text = text

    def json(self):
        return self._data


class FakeSession:
    def __init__(self, outcomes):
        self.outcomes = list(outcomes)
        self.calls = []
        self.closed = False

    def request(self, method, url, **kwargs):
        self.calls.append((method, url, kwargs))
        outcome = self.outcomes.pop(0)
        if isinstance(outcome, BaseException):
            raise outcome
        return outcome

    def close(self):
        self.closed = True


class CloudflareTempMailProviderTests(unittest.TestCase):
    def make_provider(self):
        return CloudflareTempMailProvider(
            {"api_base": "https://mail.example", "admin_password": "pw", "domain": ["example.com"]},
            {"request_timeout": 1, "wait_timeout": 1, "wait_interval": 0.2, "user_agent": "ua"},
        )

    def test_create_mailbox_retries_transient_tls_failure(self):
        provider = self.make_provider()
        fake_session = FakeSession([
            requests.exceptions.SSLError("TLS connect error"),
            FakeResponse(data={"address": "name@example.com", "jwt": "jwt-token"}),
        ])
        provider.session = fake_session

        with mock.patch("services.register.mail_provider.time.sleep"):
            mailbox = provider.create_mailbox("name")

        self.assertEqual(mailbox["address"], "name@example.com")
        self.assertEqual(mailbox["token"], "jwt-token")
        self.assertEqual(len(fake_session.calls), 2)

    def test_request_does_not_retry_invalid_domain(self):
        provider = self.make_provider()
        fake_session = FakeSession([FakeResponse(status_code=400, text="Failed to create address: Invalid domain")])
        provider.session = fake_session

        with self.assertRaisesRegex(RuntimeError, "HTTP 400"):
            provider.create_mailbox("name")

        self.assertEqual(len(fake_session.calls), 1)


class YydsMailProviderTests(unittest.TestCase):
    def make_provider(self):
        return YydsMailProvider(
            {"api_base": "https://maliapi.example/v1", "api_key": "key", "domain": ["example.com"]},
            {"request_timeout": 1, "wait_timeout": 1, "wait_interval": 0.2, "user_agent": "ua"},
        )

    def test_create_mailbox_retries_transient_tls_failure(self):
        provider = self.make_provider()
        fake_session = FakeSession([
            requests.exceptions.SSLError("unexpected eof"),
            FakeResponse(data={"data": {"address": "name@example.com", "token": "mail-token"}}),
        ])
        provider.session = fake_session

        with mock.patch("services.register.mail_provider.time.sleep"):
            mailbox = provider.create_mailbox("name")

        self.assertEqual(mailbox["address"], "name@example.com")
        self.assertEqual(mailbox["token"], "mail-token")
        self.assertEqual(len(fake_session.calls), 2)

    def test_request_does_not_retry_bad_request(self):
        provider = self.make_provider()
        fake_session = FakeSession([FakeResponse(status_code=400, text="bad domain")])
        provider.session = fake_session

        with self.assertRaisesRegex(RuntimeError, "HTTP 400"):
            provider.create_mailbox("name")

        self.assertEqual(len(fake_session.calls), 1)


if __name__ == "__main__":
    unittest.main()