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
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
|
import secrets
from urllib.parse import parse_qs, urlparse
import pytest
from httpx import AsyncClient
from jb.api.auth import SESSION_COOKIE_NAME
from jb.dependencies import get_gr_api_manager
from jb.main import app
from jb.models.auth import User
class FakeGRApiManager:
def __init__(self):
self.users: dict[str, User] = {}
self.amt_worker_ids: set[str] = set()
self.transition_calls: list[tuple[User, str]] = []
def ensure_user_exists(self, user: User) -> User:
self.users.setdefault(user.product_user_id, user)
return self.users[user.product_user_id]
def get_user(self, product_user_id: str) -> User:
return self.users[product_user_id]
def add_amt_user(self, amt_worker_id: str) -> None:
self.amt_worker_ids.add(amt_worker_id)
def transition_user_from_amt(self, user: User, amt_worker_id: str) -> User:
self.transition_calls.append((user, amt_worker_id))
if amt_worker_id not in self.amt_worker_ids:
raise ValueError(f"User {amt_worker_id} does not exist")
self.amt_worker_ids.remove(amt_worker_id)
self.users[user.product_user_id] = user
return user
@pytest.fixture
def email() -> str:
email = secrets.token_urlsafe(16) + "@gmail.com"
return email.lower()
@pytest.fixture
def fake_gr_api_manager():
manager = FakeGRApiManager()
app.dependency_overrides[get_gr_api_manager] = lambda: manager
yield manager
app.dependency_overrides.pop(get_gr_api_manager, None)
class TestAuth:
@pytest.mark.anyio
async def test_magic_link(
self,
httpxclient: AsyncClient,
fake_gr_api_manager: FakeGRApiManager,
email: str,
):
client = httpxclient
res = await client.post("/auth/magic-link/request", json={"email": email})
d = res.json()
assert res.status_code == 200
assert d["magic_link"]
token = parse_qs(urlparse(d["magic_link"]).query)["token"][0]
url = "/auth/magic-link/exchange"
body = {"token": token}
res = await client.post(url, json=body)
assert res.status_code == 204
assert client.cookies.get(SESSION_COOKIE_NAME)
res = await client.get("/auth/session")
assert res.status_code == 200
assert res.json()["email"] == email
@pytest.mark.anyio
async def test_amt_account_link(
self,
httpxclient: AsyncClient,
fake_gr_api_manager: FakeGRApiManager,
email: str,
amt_worker_id: str,
):
client = httpxclient
fake_gr_api_manager.add_amt_user(amt_worker_id)
res = await client.post(
"/auth/link-amt/request",
json={"email": email, "amt_worker_id": amt_worker_id},
)
d = res.json()
assert res.status_code == 200
assert d["magic_link"]
assert fake_gr_api_manager.transition_calls == []
token = parse_qs(urlparse(d["magic_link"]).query)["token"][0]
url = "/auth/link-amt/exchange"
body = {"token": token}
res = await client.post(url, json=body)
assert res.status_code == 204
assert client.cookies.get(SESSION_COOKIE_NAME)
assert len(fake_gr_api_manager.transition_calls) == 1
transitioned_user, transitioned_amt_worker_id = (
fake_gr_api_manager.transition_calls[0]
)
assert transitioned_user.email == email
assert transitioned_amt_worker_id == amt_worker_id
assert amt_worker_id not in fake_gr_api_manager.amt_worker_ids
res = await client.get("/auth/session")
assert res.status_code == 200
assert res.json()["email"] == email
# Make sure we can't do it again
res = await client.post(url, json=body)
assert res.status_code == 401
assert len(fake_gr_api_manager.transition_calls) == 1
|