Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 2 additions & 7 deletions src/py/mat3ra/api_client/client.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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"]

Expand Down
18 changes: 4 additions & 14 deletions src/py/mat3ra/api_client/endpoints/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import json # noqa: F401

from ..models import AuthContext
from ..utils.http import Connection


Expand All @@ -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.
Expand All @@ -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}
5 changes: 5 additions & 0 deletions src/py/mat3ra/api_client/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
25 changes: 25 additions & 0 deletions tests/py/unit/test_auth_context.py
Original file line number Diff line number Diff line change
@@ -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
38 changes: 24 additions & 14 deletions tests/py/unit/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}}}}

Expand Down Expand Up @@ -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):
Expand Down
Loading