diff --git a/src/py/mat3ra/api_client/client.py b/src/py/mat3ra/api_client/client.py index 1299e27..060acdd 100644 --- a/src/py/mat3ra/api_client/client.py +++ b/src/py/mat3ra/api_client/client.py @@ -1,11 +1,10 @@ -import os import re from typing import Any, List, Optional, Tuple import requests from pydantic import BaseModel, ConfigDict -from .constants import ACCESS_TOKEN_ENV_VAR, _build_base_url +from .constants import _build_base_url from .endpoints.bank_materials import BankMaterialEndpoints from .endpoints.bank_workflows import BankWorkflowEndpoints from .endpoints.clusters import ClustersEndpoint @@ -128,12 +127,8 @@ def authenticate( ) def _fetch_data(self) -> dict: - access_token = self.auth.access_token or os.environ.get(ACCESS_TOKEN_ENV_VAR) - if not access_token: - raise ValueError("Access token is required to fetch user data") - url = _build_base_url(self.host, self.port, self.secure, "/api/v1/users/me") - response = requests.get(url, headers={"Authorization": f"Bearer {access_token}"}, timeout=30) + response = requests.get(url, headers=self.auth.get_headers(), timeout=30) response.raise_for_status() return response.json()["data"] diff --git a/src/py/mat3ra/api_client/endpoints/__init__.py b/src/py/mat3ra/api_client/endpoints/__init__.py index 5eae989..a568656 100644 --- a/src/py/mat3ra/api_client/endpoints/__init__.py +++ b/src/py/mat3ra/api_client/endpoints/__init__.py @@ -1,5 +1,6 @@ import json # noqa: F401 +from ..models import AuthContext from ..utils.http import Connection @@ -21,12 +22,6 @@ def __init__(self, host, port, version="2018-10-1", secure=True, **kwargs): self._auth = kwargs.get("auth") self.conn = Connection(host, port, version=version, secure=secure, **kwargs) - def _get_bearer_headers(self): - access_token = getattr(self._auth, "access_token", None) - if access_token: - return {"Authorization": f"Bearer {access_token}"} - return {} - def request(self, method, endpoint_path, params=None, data=None, headers=None): """ Sends an HTTP request with given params, headers and data to the given endpoint. @@ -41,18 +36,13 @@ def request(self, method, endpoint_path, params=None, data=None, headers=None): Returns: json: response """ - request_headers = dict(headers or {}) - bearer_headers = self._get_bearer_headers() - if bearer_headers: - request_headers.update(bearer_headers) - request_headers.pop("X-Account-Id", None) - request_headers.pop("X-Auth-Token", None) with self.conn: - self.conn.request(method, endpoint_path, params, data, request_headers or None) + self.conn.request(method, endpoint_path, params, data, headers) response = self.conn.json() if response["status"] != "success": raise BaseException(response["data"]["message"]) return response["data"] def get_headers(self, account_id, auth_token, content_type="application/json"): - return {"X-Account-Id": account_id, "X-Auth-Token": auth_token, "Content-Type": content_type} + auth = self._auth or AuthContext(account_id=account_id, auth_token=auth_token) + return {**auth.get_headers(), "Content-Type": content_type} diff --git a/src/py/mat3ra/api_client/models.py b/src/py/mat3ra/api_client/models.py index 651d4ab..c26915b 100644 --- a/src/py/mat3ra/api_client/models.py +++ b/src/py/mat3ra/api_client/models.py @@ -11,6 +11,11 @@ class AuthContext(BaseModel): account_id: Optional[str] = None auth_token: Optional[str] = None + def get_headers(self) -> dict: + if self.access_token: + return {"Authorization": f"Bearer {self.access_token}"} + return {"X-Account-Id": self.account_id, "X-Auth-Token": self.auth_token} + class APIEnv(BaseModel): host: str = Field(default="platform.mat3ra.com", validation_alias="API_HOST") diff --git a/tests/py/unit/test_auth_context.py b/tests/py/unit/test_auth_context.py new file mode 100644 index 0000000..669f278 --- /dev/null +++ b/tests/py/unit/test_auth_context.py @@ -0,0 +1,25 @@ +import pytest +from mat3ra.api_client import AuthContext + +OIDC_ACCESS_TOKEN = "oidc-access-token" +ACCOUNT_ID = "ubxMkAyx37Rjn8qK9" +AUTH_TOKEN = "legacy-auth-token" + +OIDC_AUTH = {"access_token": OIDC_ACCESS_TOKEN} +API_TOKEN_AUTH = {"account_id": ACCOUNT_ID, "auth_token": AUTH_TOKEN} +OIDC_AND_API_TOKEN_AUTH = OIDC_AUTH | API_TOKEN_AUTH + +BEARER_HEADERS = {"Authorization": f"Bearer {OIDC_ACCESS_TOKEN}"} +API_TOKEN_HEADERS = {"X-Account-Id": ACCOUNT_ID, "X-Auth-Token": AUTH_TOKEN} + + +@pytest.mark.parametrize( + "auth, expected_headers", + [ + (OIDC_AUTH, BEARER_HEADERS), + (API_TOKEN_AUTH, API_TOKEN_HEADERS), + (OIDC_AND_API_TOKEN_AUTH, BEARER_HEADERS), + ], +) +def test_get_headers(auth, expected_headers): + assert AuthContext(**auth).get_headers() == expected_headers diff --git a/tests/py/unit/test_client.py b/tests/py/unit/test_client.py index b19aa41..2e6b65d 100644 --- a/tests/py/unit/test_client.py +++ b/tests/py/unit/test_client.py @@ -14,6 +14,14 @@ AUTH_TOKEN = "legacy-auth-token" ACCOUNT_ID = "ubxMkAyx37Rjn8qK9" +AUTH_ENV_AND_USERS_ME_HEADERS = [ + ({"OIDC_ACCESS_TOKEN": OIDC_ACCESS_TOKEN}, {"Authorization": f"Bearer {OIDC_ACCESS_TOKEN}"}), + ( + {"ACCOUNT_ID": ACCOUNT_ID, "AUTH_TOKEN": AUTH_TOKEN}, + {"X-Account-Id": ACCOUNT_ID, "X-Auth-Token": AUTH_TOKEN}, + ), +] + ME_ACCOUNT_ID = "my-account-id" USERS_ME_RESPONSE = {"data": {"user": {"entity": {"defaultAccountId": ME_ACCOUNT_ID}}}} @@ -104,20 +112,22 @@ def test_my_account_id_fetches_and_caches(self, mock_get): @mock.patch("requests.get") def test_list_accounts(self, mock_get): - env = self._base_env() | {"OIDC_ACCESS_TOKEN": OIDC_ACCESS_TOKEN} - with mock.patch.dict("os.environ", env, clear=True): - self._mock_users_me(mock_get, ACCOUNTS_RESPONSE) - client = APIClient.authenticate() - accounts = client.list_accounts() - - self.assertEqual(len(accounts), 3) - self.assertEqual(accounts[0]["_id"], "user-acc-1") - self.assertEqual(accounts[0]["name"], "John Doe") - self.assertEqual(accounts[0]["type"], "personal") - self.assertTrue(accounts[0]["isDefault"]) - self.assertEqual(accounts[1]["_id"], "org-acc-1") - self.assertEqual(accounts[1]["name"], "Acme Corp") - self.assertEqual(accounts[1]["type"], "enterprise") + for auth_env, expected_headers in AUTH_ENV_AND_USERS_ME_HEADERS: + with self.subTest(auth_env=auth_env), mock.patch.dict( + "os.environ", self._base_env() | auth_env, clear=True + ): + self._mock_users_me(mock_get, ACCOUNTS_RESPONSE) + accounts = APIClient.authenticate().list_accounts() + + self.assertEqual(mock_get.call_args[1]["headers"], expected_headers) + self.assertEqual(len(accounts), 3) + self.assertEqual(accounts[0]["_id"], "user-acc-1") + self.assertEqual(accounts[0]["name"], "John Doe") + self.assertEqual(accounts[0]["type"], "personal") + self.assertTrue(accounts[0]["isDefault"]) + self.assertEqual(accounts[1]["_id"], "org-acc-1") + self.assertEqual(accounts[1]["name"], "Acme Corp") + self.assertEqual(accounts[1]["type"], "enterprise") @mock.patch("requests.get") def test_get_account(self, mock_get):