diff options
| author | Max Nanis | 2026-09-13 19:03:55 +0000 |
|---|---|---|
| committer | Max Nanis | 2026-09-13 19:03:55 +0000 |
| commit | 8fbe8d439b418796932aaa88297e006c714e4d13 (patch) | |
| tree | 6c800476edc1e771fc570b559d483df8d59d4324 /tests/http/test_auth.py | |
| parent | 12f6fee851b68e86af658dfa17e4a0daed457dd1 (diff) | |
| parent | 2c94f248d2438071a918fa9a30bf114ef9aa29b4 (diff) | |
| download | amt-jb-8fbe8d439b418796932aaa88297e006c714e4d13.tar.gz amt-jb-8fbe8d439b418796932aaa88297e006c714e4d13.zip | |
Merges pull request #3
Off of Amazon!!!
Diffstat (limited to 'tests/http/test_auth.py')
| -rw-r--r-- | tests/http/test_auth.py | 120 |
1 files changed, 120 insertions, 0 deletions
diff --git a/tests/http/test_auth.py b/tests/http/test_auth.py new file mode 100644 index 0000000..1fa1335 --- /dev/null +++ b/tests/http/test_auth.py @@ -0,0 +1,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 |
