diff --git a/src/py/mat3ra/api_client/endpoints/__init__.py b/src/py/mat3ra/api_client/endpoints/__init__.py index a568656..3316f9f 100644 --- a/src/py/mat3ra/api_client/endpoints/__init__.py +++ b/src/py/mat3ra/api_client/endpoints/__init__.py @@ -19,7 +19,7 @@ class BaseEndpoint(object): """ def __init__(self, host, port, version="2018-10-1", secure=True, **kwargs): - self._auth = kwargs.get("auth") + self.auth = kwargs.get("auth") self.conn = Connection(host, port, version=version, secure=secure, **kwargs) def request(self, method, endpoint_path, params=None, data=None, headers=None): @@ -36,6 +36,8 @@ def request(self, method, endpoint_path, params=None, data=None, headers=None): Returns: json: response """ + if headers and self.auth: + headers = {**headers, **self.auth.get_headers()} with self.conn: self.conn.request(method, endpoint_path, params, data, headers) response = self.conn.json() @@ -44,5 +46,5 @@ def request(self, method, endpoint_path, params=None, data=None, headers=None): return response["data"] def get_headers(self, account_id, auth_token, content_type="application/json"): - auth = self._auth or AuthContext(account_id=account_id, auth_token=auth_token) + 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/utils/http.py b/src/py/mat3ra/api_client/utils/http.py index c91de1e..7ee9cc5 100644 --- a/src/py/mat3ra/api_client/utils/http.py +++ b/src/py/mat3ra/api_client/utils/http.py @@ -40,7 +40,9 @@ def request(self, method, url, params=None, data=None, headers=None): data (dict): the body to attach to the request. params (dict): URL parameters to append to the URL. """ - self.response = self.session.request(method=method.lower(), url=url, params=params, data=data, headers=headers) + self.response = self.session.request( + method=method.lower(), url=url, params=params, data=data, headers=headers, timeout=self.session.timeout + ) try: self.response.raise_for_status() except requests.HTTPError: diff --git a/tests/py/unit/test_client.py b/tests/py/unit/test_client.py index 2e6b65d..cac678b 100644 --- a/tests/py/unit/test_client.py +++ b/tests/py/unit/test_client.py @@ -1,7 +1,7 @@ import os from unittest import mock -from mat3ra.api_client import APIClient +from mat3ra.api_client import APIClient, AuthContext from tests.py.unit import EndpointBaseUnitTest @@ -11,6 +11,7 @@ API_SECURE_FALSE = "false" OIDC_ACCESS_TOKEN = "oidc-access-token" +NEW_OIDC_ACCESS_TOKEN = "new-oidc-access-token" AUTH_TOKEN = "legacy-auth-token" ACCOUNT_ID = "ubxMkAyx37Rjn8qK9" @@ -152,3 +153,15 @@ def test_my_organization(self, mock_get): org = client.my_organization self.assertEqual(org.id, "org-acc-1") self.assertEqual(org.name, "Acme Corp") + + @mock.patch("requests.sessions.Session.request") + def test_endpoint_request_sends_current_token_and_timeout(self, mock_request): + auth = AuthContext(access_token=OIDC_ACCESS_TOKEN) + client = APIClient( + host=API_HOST, port=API_PORT, version=API_VERSION, secure=False, auth=auth, timeout_seconds=5 + ) + client.auth.access_token = NEW_OIDC_ACCESS_TOKEN + mock_request.return_value.json.return_value = {"status": "success", "data": []} + client.jobs.list() + self.assertEqual(mock_request.call_args[1]["headers"]["Authorization"], f"Bearer {NEW_OIDC_ACCESS_TOKEN}") + self.assertEqual(mock_request.call_args[1]["timeout"], 5)