aboutsummaryrefslogtreecommitdiff
path: root/tests/http/test_auth.py
blob: 1fa1335b49bfcb7defbc445d702689938729c3ad (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
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:
    return "unittest@generalresearch.com"


@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