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
6 changes: 4 additions & 2 deletions src/py/mat3ra/api_client/endpoints/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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()
Expand All @@ -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}
4 changes: 3 additions & 1 deletion src/py/mat3ra/api_client/utils/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
15 changes: 14 additions & 1 deletion tests/py/unit/test_client.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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"

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