aboutsummaryrefslogtreecommitdiff
path: root/tests/http/test_auth.py
blob: 02ac88aad8a67f270f315d7b764e512f3f0d01e9 (plain)
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
123
124
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