aboutsummaryrefslogtreecommitdiff
path: root/jb/managers/gr_api.py
blob: 8ffeaf0f96a9930c88d4170aa8cf1016e1b1ee60 (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
"""Client for General Research's product-user API."""

from typing import Any

import requests

from jb.models.auth import User


class GRApiError(RuntimeError):
    """The General Research API could not satisfy a request."""


class GRApiManager:
    def __init__(
        self,
        base_url: str,
        token: str,
        product_id: str,
        timeout: float = 5.0,
        session: requests.Session | None = None,
    ) -> None:
        if not token:
            raise ValueError("General Research API token must not be empty")

        self.base_url = base_url.rstrip("/")
        self.product_id = product_id
        self.timeout = timeout
        self.session = session or requests.Session()
        self.session.headers.update(
            {
                "Authorization": token,
                "Accept": "application/json",
                "Content-Type": "application/json",
            }
        )

    def _request(self, method: str, url: str, **kwargs: Any) -> requests.Response:
        try:
            response = self.session.request(method, url, timeout=self.timeout, **kwargs)
            response.raise_for_status()
            return response
        except requests.RequestException as exc:
            raise GRApiError(
                f"General Research API request failed: {method} {url}"
            ) from exc

    def _parse_user_response(self, res: dict) -> User:
        return User.model_validate(
            {
                "product_user_id": res["product_user_id"],
                "email": res["metadata"].get("email_address"),
                "display_name": res["metadata"].get("display_name"),
            }
        )

    def ensure_user_exists(self, user: User) -> User:
        """Idempotently create the product user if it does not already exist."""
        url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/"
        res = self._request("PUT", url).json()
        if res["metadata"].get("email_address") is None:
            self.set_user_email(user)
            return self.get_user(user.product_user_id)
        user_thl = self._parse_user_response(res)
        if user_thl.email != user.email:
            raise ValueError(
                f"user {user.product_user_id} already exists with email {user_thl.email}"
            )
        return user_thl

    def get_user(self, product_user_id: str) -> User:
        """Retrieve the user's email address and display name."""
        url = f"{self.base_url}/{self.product_id}/user/{product_user_id}/"
        res = self._request("GET", url).json()
        return self._parse_user_response(res)

    def get_user_by_email(self, email: str) -> User:
        user = User.model_validate({"email": email})
        return self.get_user(user.product_user_id)

    def set_user_email(self, user: User) -> None:
        """This should only be called once per user upon account creation.
        A user cannot change their email address."""
        url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/metadata/"
        res = self._request(
            "PATCH",
            url,
            json={"email_address": str(user.email)},
        ).json()

    def set_user_display_name(self, user: User) -> User:
        """Can be called as many times as needed. Multiple
        users can have the same display name."""
        if user.display_name is None:
            return user
        url = f"{self.base_url}/{self.product_id}/user/{user.product_user_id}/metadata/"
        res = self._request(
            "PATCH",
            url,
            json={"display_name": user.display_name},
        ).json()
        return self.get_user(user.product_user_id)

    def transition_product_user_id(self, user: User, amt_worker_id: str) -> User:
        """This should only be called once upon transition from an
        AMT account to a General Research account."""
        url = f"{self.base_url}/{self.product_id}/user/{amt_worker_id}/"
        res = self._request(
            "PATCH", url, json={"product_user_id": user.product_user_id}
        )
        self.set_user_email(user)
        return self.get_user(user.product_user_id)