diff --git a/config.yml b/config.yml index 5a5c5477d..7b88d0f05 100644 --- a/config.yml +++ b/config.yml @@ -66,7 +66,7 @@ notebooks: - ipywidgets - tabulate==0.9.0 - mat3ra-utils - - mat3ra-api-client + - mat3ra-api-client>=2026.10.8.post0 - name: dataframe packages_pyodide: - jinja2 @@ -79,7 +79,7 @@ notebooks: - emfs:/drive/packages/pydantic-2.7.1-py3-none-any.whl - mat3ra-standata - mat3ra-ide - - mat3ra-api-client + - mat3ra-api-client>=2026.10.8.post0 - requests - mat3ra-ade - mat3ra-mode diff --git a/pyproject.toml b/pyproject.toml index 0f4f23014..2840c8d43 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -24,7 +24,7 @@ utils_standata = [ "mat3ra-standata", ] api = [ - "mat3ra-api-client", + "mat3ra-api-client>=2026.10.8.post0", ] materials = [ # ase >=3.25.0 is required for supercell to be generated, diff --git a/src/py/mat3ra/notebooks_utils/api/job.py b/src/py/mat3ra/notebooks_utils/api/job.py index 991ca56fe..ad8d96d73 100644 --- a/src/py/mat3ra/notebooks_utils/api/job.py +++ b/src/py/mat3ra/notebooks_utils/api/job.py @@ -1,20 +1,22 @@ +import asyncio import datetime from collections import Counter -from typing import List +from typing import Any, List +import requests from mat3ra.api_client import JobEndpoints from mat3ra.utils.extra.tabulate import pretty_print -from ..core.entity.job.api import create_job, get_jobs_statuses_by_ids, save_files, submit_jobs +from ..core.entity.job.api import create_job, get_jobs_statuses_by_ids_async, save_files, submit_jobs from ..pyodide.runtime import interruptible_polling_loop # Contains no external dependencies, only uses the API client, # so can be used in both regular Python and Pyodide environments. -# In regular Python, the interruptible_polling_loop does not actually do anything, -# so the function behaves like a normal polling loop. @interruptible_polling_loop() -def wait_for_jobs_to_finish_async(endpoint: JobEndpoints, job_ids: List[str]) -> bool: +async def wait_for_jobs_to_finish_async( + endpoint: JobEndpoints, job_ids: List[str], *, abort_signal: Any = None +) -> bool: """ Waits for jobs to finish and prints their statuses. A job is considered finished if it is not in "pre-submission", "submitted", or "active" status. @@ -22,8 +24,15 @@ def wait_for_jobs_to_finish_async(endpoint: JobEndpoints, job_ids: List[str]) -> Args: endpoint (JobEndpoints): Job endpoint object from the Exabyte API Client job_ids (list): list of job IDs to wait for + abort_signal: JS AbortSignal of the Abort control, passed in by the polling loop """ - statuses = get_jobs_statuses_by_ids(endpoint, job_ids) + try: + statuses = await get_jobs_statuses_by_ids_async(endpoint, job_ids, abort_signal=abort_signal) + except (asyncio.TimeoutError, OSError) as error: + if isinstance(error, requests.HTTPError) and error.response.status_code < 500: + raise + print(f"status check failed: {error!r}, retrying") + return True counts = Counter(statuses) headers = ["TIME", "SUBMITTED-JOBS", "ACTIVE-JOBS", "FINISHED-JOBS", "ERRORED-JOBS"] now = datetime.datetime.now().strftime("%Y-%m-%d-%H:%M:%S") @@ -42,7 +51,6 @@ def wait_for_jobs_to_finish_async(endpoint: JobEndpoints, job_ids: List[str]) -> __all__ = [ "create_job", - "get_jobs_statuses_by_ids", "save_files", "submit_jobs", "wait_for_jobs_to_finish_async", diff --git a/src/py/mat3ra/notebooks_utils/auth.py b/src/py/mat3ra/notebooks_utils/auth.py index 433f36490..ec5fd1f9e 100644 --- a/src/py/mat3ra/notebooks_utils/auth.py +++ b/src/py/mat3ra/notebooks_utils/auth.py @@ -13,26 +13,37 @@ import inspect import os -from mat3ra.api_client import ACCESS_TOKEN_ENV_VAR +import requests +from mat3ra.api_client import ACCESS_TOKEN_ENV_VAR, APIClient, AuthContext from .core.api.auth import authenticate_oidc, get_oidc_base_url, store_token_data_in_environment from .io import get_data from .ipython.ui import show_device_flow_popup from .primitive.environment import is_pyodide_environment from .pyodide.api.auth import authenticate_jupyterlite -from .token_store import load_token, save_token +from .token_store import delete_token, load_token, save_token REFRESH_TOKEN_ENV_VAR = "OIDC_REFRESH_TOKEN" async def _authenticate_oidc_with_cache(force=False): oidc_url = get_oidc_base_url() - cached = None if force else await load_token(oidc_url) - - if cached: - store_token_data_in_environment(cached) - return + access_token = os.environ.get(ACCESS_TOKEN_ENV_VAR) + token_data = {"access_token": access_token} if access_token else await load_token(oidc_url) + if token_data and not force: + try: + APIClient.authenticate(access_token=token_data["access_token"]).list_accounts() + store_token_data_in_environment(token_data) + return + except KeyError: + pass + except requests.HTTPError as error: + if error.response.status_code != 401: + raise + + await delete_token(oidc_url) + os.environ.pop(ACCESS_TOKEN_ENV_VAR, None) token_data = await authenticate_oidc(show_popup=show_device_flow_popup) await save_token(oidc_url, token_data) @@ -61,5 +72,13 @@ async def authenticate(force=False, globals_dict=None): if data_from_host: await authenticate_jupyterlite(data_from_host) - elif ACCESS_TOKEN_ENV_VAR not in os.environ or force: + else: await _authenticate_oidc_with_cache(force) + + +async def reauthenticate(auth_context: AuthContext) -> None: + """ + Replaces an access token the platform rejected with a new one from the device login, set on `auth_context`. + """ + await _authenticate_oidc_with_cache(force=True) + auth_context.access_token = os.environ[ACCESS_TOKEN_ENV_VAR] diff --git a/src/py/mat3ra/notebooks_utils/core/api/auth.py b/src/py/mat3ra/notebooks_utils/core/api/auth.py index 91d21f7cc..1b88b5e21 100644 --- a/src/py/mat3ra/notebooks_utils/core/api/auth.py +++ b/src/py/mat3ra/notebooks_utils/core/api/auth.py @@ -1,12 +1,23 @@ import asyncio import os import time -from typing import Callable, Optional +import urllib.parse +from typing import Any, Awaitable, Callable, Optional import requests from mat3ra.api_client import ACCESS_TOKEN_ENV_VAR, CLIENT_ID, SCOPE, APIEnv, build_oidc_base_url +from ...primitive.environment import is_pyodide_environment +from ...pyodide.runtime import run_interruptible_loop_async + +try: + from pyodide.http import pyfetch # type: ignore +except ImportError: + pyfetch = None + REFRESH_TOKEN_ENV_VAR = "OIDC_REFRESH_TOKEN" +TOKEN_REQUEST_TIMEOUT_SECONDS = 10 +FORM_HEADERS = {"Content-Type": "application/x-www-form-urlencoded"} def get_oidc_base_url() -> str: @@ -50,6 +61,20 @@ def store_token_data_in_environment(token_data: dict) -> None: os.environ[REFRESH_TOKEN_ENV_VAR] = token_data["refresh_token"] +def _request_token_data(token_url: str, form_data: dict) -> tuple: + response = requests.post(token_url, data=form_data, headers=FORM_HEADERS, timeout=TOKEN_REQUEST_TIMEOUT_SECONDS) + return response.status_code, response.json() if response.status_code < 500 else {} + + +async def _request_token_data_with_fetch(token_url: str, form_data: dict, abort_signal: Any) -> tuple: + """ + `_request_token_data` through the browser's fetch, which leaves the event loop free while the request is in flight. + """ + body = urllib.parse.urlencode(form_data, doseq=True) + response = await pyfetch(token_url, method="POST", body=body, headers=FORM_HEADERS, signal=abort_signal) + return response.status, await response.json() if response.status < 500 else {} + + async def _poll_for_token_data( oidc_base_url: str, client_id: str, @@ -57,25 +82,37 @@ async def _poll_for_token_data( polling_interval_seconds: int, expires_in_seconds: int, ) -> dict: + """Polls for the token until the device login is confirmed; in pyodide, ESC stops it at once.""" + token_url = f"{oidc_base_url}/token" + form_data = { + "grant_type": "urn:ietf:params:oauth:grant-type:device_code", + "device_code": device_code, + "client_id": client_id, + "redirect_uris": [], + "response_types": [], + "token_endpoint_auth_method": "none", + } deadline_seconds = time.time() + expires_in_seconds - while time.time() < deadline_seconds: - token_response = requests.post( - f"{oidc_base_url}/token", - data={ - "grant_type": "urn:ietf:params:oauth:grant-type:device_code", - "device_code": device_code, - "client_id": client_id, - "redirect_uris": [], - "response_types": [], - "token_endpoint_auth_method": "none", - }, - headers={"Content-Type": "application/x-www-form-urlencoded"}, - timeout=10, - ) - if token_response.status_code == 200: - return token_response.json() - await asyncio.sleep(polling_interval_seconds) - raise Exception("Timeout waiting for authorization.") + token_data: dict = {} + + def request_token_data(abort_signal: Any) -> Awaitable[tuple]: + if is_pyodide_environment(): + return _request_token_data_with_fetch(token_url, form_data, abort_signal) + return asyncio.get_running_loop().run_in_executor(None, _request_token_data, token_url, form_data) + + async def poll_step(abort_signal: Any) -> bool: + if time.time() >= deadline_seconds: + raise Exception("Timeout waiting for authorization.") + status, response_data = await asyncio.wait_for(request_token_data(abort_signal), TOKEN_REQUEST_TIMEOUT_SECONDS) + if status != 200 and status < 500 and response_data.get("error") not in ("authorization_pending", "slow_down"): + raise Exception(f"Device login failed ({status}): {response_data.get('error')}.") + token_data.update(response_data if status == 200 else {}) + return not token_data + + await run_interruptible_loop_async( + poll_step, polling_interval_seconds, show_button=False, abort_hint_text="Press ESC to cancel" + ) + return token_data async def authenticate_oidc( diff --git a/src/py/mat3ra/notebooks_utils/core/entity/job/api.py b/src/py/mat3ra/notebooks_utils/core/entity/job/api.py index 363c5b2f9..ca8b7440f 100644 --- a/src/py/mat3ra/notebooks_utils/core/entity/job/api.py +++ b/src/py/mat3ra/notebooks_utils/core/entity/job/api.py @@ -1,10 +1,23 @@ +import asyncio +import json import re +import urllib.parse import urllib.request -from typing import Any, Dict, Iterable, List, Optional, Union +from typing import Any, Awaitable, Dict, Iterable, List, Optional, Union +import requests from mat3ra.api_client import APIClient, JobEndpoints +from ....auth import reauthenticate +from ....primitive.environment import is_pyodide_environment + +try: + from pyodide.http import pyfetch # type: ignore +except ImportError: + pyfetch = None + MATERIALS_SET_ENTITY_CLASS = "Material" +DEFAULT_STATUS_TIMEOUT_SECONDS = 30 def save_files(job_id: str, job_endpoint: JobEndpoints, filename_on_cloud: str, filename_on_disk: str) -> None: @@ -25,18 +38,55 @@ def save_files(job_id: str, job_endpoint: JobEndpoints, filename_on_cloud: str, outp.write(server_response.read()) -def get_jobs_statuses_by_ids(endpoint: JobEndpoints, job_ids: List[str]) -> List[str]: +async def _list_jobs_with_fetch(endpoint: JobEndpoints, query: dict, projection: dict, abort_signal: Any) -> List[dict]: + """ + `endpoint.list` through the browser's fetch, which leaves the event loop free while the request is in flight. + Raises `requests.HTTPError` on an error status. + """ + parameters = urllib.parse.urlencode({"query": json.dumps(query), "projection": json.dumps(projection)}) + url = urllib.parse.urljoin(endpoint.conn.preamble, f"{endpoint.name}?{parameters}") + response = await pyfetch(url, headers={**endpoint.headers, **endpoint.auth.get_headers()}, signal=abort_signal) + if not response.ok: + error_response = requests.Response() + error_response.status_code = response.status + raise requests.HTTPError(f"Error {response.status}.", response=error_response) + return (await response.json())["data"] + + +async def get_jobs_statuses_by_ids_async( + endpoint: JobEndpoints, + job_ids: List[str], + timeout: float = DEFAULT_STATUS_TIMEOUT_SECONDS, + abort_signal: Any = None, +) -> List[str]: """ - Gets jobs statues by their IDs. + Gets jobs statuses by their IDs without blocking the event loop: through the browser's fetch in pyodide, + in a worker thread otherwise. A rejected access token (401) is replaced through the device login once. Args: endpoint (JobEndpoints): Job endpoint object from the Exabyte API Client job_ids (list): list of job IDs to get the status for + timeout (float): seconds to wait for the response before raising asyncio.TimeoutError + abort_signal: JS AbortSignal that aborts the fetch in pyodide Returns: list: list of job statuses """ - jobs = endpoint.list({"_id": {"$in": job_ids}}, {"fields": {"status": 1}}) + query = {"_id": {"$in": job_ids}} + projection = {"fields": {"status": 1}} + + def request_jobs() -> Awaitable[List[dict]]: + if is_pyodide_environment(): + return _list_jobs_with_fetch(endpoint, query, projection, abort_signal) + return asyncio.get_running_loop().run_in_executor(None, endpoint.list, query, projection) + + try: + jobs = await asyncio.wait_for(request_jobs(), timeout) + except requests.HTTPError as error: + if error.response.status_code != 401 or not endpoint.auth.access_token: + raise + await reauthenticate(endpoint.auth) + jobs = await asyncio.wait_for(request_jobs(), timeout) return [job["status"] for job in jobs] diff --git a/src/py/mat3ra/notebooks_utils/pyodide/runtime.py b/src/py/mat3ra/notebooks_utils/pyodide/runtime.py index efac1ffa6..2e60a8906 100644 --- a/src/py/mat3ra/notebooks_utils/pyodide/runtime.py +++ b/src/py/mat3ra/notebooks_utils/pyodide/runtime.py @@ -14,14 +14,18 @@ HTML = None display = None +ABORT_CHANNEL_NAME = f"mat3ra_abort_{uuid.uuid4().hex}" + class UserAbortError(RuntimeError): pass def display_abort_controls_in_current_cell_output( - channel_name: str = "mat3ra_abort_channel", + channel_name: str = ABORT_CHANNEL_NAME, abort_button_text: str = "Stop polling", + abort_hint_text: str = "Press ESC to abort", + show_button: bool = True, ) -> None: """ Shows: @@ -34,20 +38,24 @@ def display_abort_controls_in_current_cell_output( return element_id = f"abort_controls_{uuid.uuid4().hex}" - - display( - HTML( - f""" -
+ button_html = f""" + >{abort_button_text}""" - Press ESC to abort + display( + HTML( + f""" +
+ {button_html if show_button else ""} + {abort_hint_text}
@@ -56,13 +64,18 @@ def display_abort_controls_in_current_cell_output( const channelName = {channel_name!r}; // Install ESC broadcaster once per page - if (!window.__mat3raEscapeAbortInstalled) {{ - window.__mat3raEscapeAbortInstalled = true; - const escChannel = new BroadcastChannel(channelName); + if (!window.__mat3raEscapeAbortByPanelInstalled) {{ + window.__mat3raEscapeAbortByPanelInstalled = true; document.addEventListener("keydown", (event) => {{ - if (event.key === "Escape") {{ - escChannel.postMessage({{ type: "abort", source: "escape" }}); - }} + if (event.key !== "Escape") return; + const notebookPanel = document.activeElement?.closest(".jp-NotebookPanel"); + const abortHints = (notebookPanel || document).querySelectorAll("[data-mat3ra-abort-channel]"); + if (!notebookPanel && abortHints.length !== 1) return; + abortHints.forEach((abortHint) => {{ + const escapeChannel = new BroadcastChannel(abortHint.dataset.mat3raAbortChannel); + escapeChannel.postMessage({{ type: "abort", source: "escape" }}); + escapeChannel.close(); + }}); }}, true); }} @@ -89,17 +102,22 @@ def display_abort_controls_in_current_cell_output( class BroadcastChannelAbortController(BaseModel): """ WebWorker-side receiver. Works only in pyodide (emscripten). - In regular Python: start() does nothing and is_aborted stays False. + In regular Python: start() does nothing, `is_aborted` stays False and `fetch_abort_signal` stays None. """ - channel_name: str = "mat3ra_abort_channel" + channel_name: str = ABORT_CHANNEL_NAME is_aborted: bool = False def model_post_init(self, __context: Any) -> None: self._broadcast_channel = None self._on_message_proxy = None + self._fetch_abort_controller = None - def start(self) -> None: + @property + def fetch_abort_signal(self) -> Any: + return getattr(self._fetch_abort_controller, "signal", None) + + def start(self, task: "asyncio.Task[Any]") -> None: if ENVIRONMENT != EnvironmentsEnum.PYODIDE: return if self._broadcast_channel is not None: @@ -109,11 +127,14 @@ def start(self) -> None: from pyodide.ffi import create_proxy # type: ignore self._broadcast_channel = js.BroadcastChannel.new(self.channel_name) + self._fetch_abort_controller = js.AbortController.new() def on_message(event) -> None: message = getattr(event, "data", None) if message and getattr(message, "type", None) == "abort": self.is_aborted = True + self._fetch_abort_controller.abort() # type: ignore + task.cancel() self._on_message_proxy = create_proxy(on_message) self._broadcast_channel.onmessage = self._on_message_proxy # type: ignore @@ -131,43 +152,42 @@ def stop(self) -> None: async def run_interruptible_loop_async( - loop_body: Callable[[], Awaitable[bool]], + loop_body: Callable[[Any], Awaitable[bool]], poll_interval_seconds: float, *, - channel_name: str = "mat3ra_abort_channel", - check_interval_seconds: float = 0.05, + channel_name: str = ABORT_CHANNEL_NAME, show_controls: bool = True, + show_button: bool = True, + abort_hint_text: str = "Press ESC to abort", ) -> None: """ Wraps an async loop around a "poll" function that returns True to continue, False to stop. - loop_body(): - - do one "poll" iteration + loop_body(abort_signal): + - do one "poll" iteration; `abort_signal` aborts its fetch (pyodide), None in regular Python - return True to keep looping, False to stop normally - Between iterations we sleep in small slices so: - - pyodide: ESC/button can be received and stop the loop - - regular Python: yields control (Ctrl+C/Stop works where supported) + pyodide: ESC/button cancels the task running the loop, during a poll or the sleep, raising UserAbortError. + regular Python: Ctrl+C/Stop interrupts the kernel. """ broadcast_channel_abort_controller = BroadcastChannelAbortController(channel_name=channel_name) - broadcast_channel_abort_controller.start() + broadcast_channel_abort_controller.start(asyncio.current_task()) # type: ignore if show_controls and ENVIRONMENT == EnvironmentsEnum.PYODIDE: - display_abort_controls_in_current_cell_output(channel_name=channel_name, abort_button_text="Abort") + display_abort_controls_in_current_cell_output( + channel_name=channel_name, + abort_button_text="Abort", + abort_hint_text=abort_hint_text, + show_button=show_button, + ) try: - while True: - should_continue = await loop_body() - if not should_continue: - return - - remaining_seconds = float(poll_interval_seconds) - while remaining_seconds > 0: - if broadcast_channel_abort_controller.is_aborted: - raise UserAbortError("Aborted by user.") - await asyncio.sleep(min(check_interval_seconds, remaining_seconds)) - remaining_seconds -= check_interval_seconds - + while await loop_body(broadcast_channel_abort_controller.fetch_abort_signal): + if broadcast_channel_abort_controller.is_aborted: + raise UserAbortError("Aborted by user.") + await asyncio.sleep(poll_interval_seconds) + except asyncio.CancelledError: + raise UserAbortError("Aborted by user.") from None finally: broadcast_channel_abort_controller.stop() @@ -176,13 +196,13 @@ def interruptible_polling_loop( poll_interval_kwarg_name: str = "poll_interval", *, default_poll_interval_seconds: float = 10.0, - channel_name: str = "mat3ra_abort_channel", - check_interval_seconds: float = 0.05, + channel_name: str = ABORT_CHANNEL_NAME, show_controls: bool = True, ): """ - Turns a poll-step function into an async loop. Wrapped fn returns True to continue, False to stop. - Sleeps in small slices so ESC/Abort (notebooks) or Ctrl+C can raise UserAbortError. + Turns a poll-step function into an async loop. Wrapped fn takes `abort_signal` for its request and returns + True to continue, False to stop. + ESC/Abort (notebooks) raises UserAbortError at once, during a poll or the sleep; Ctrl+C natively. Poll interval: kwarg poll_interval_kwarg_name, else default_poll_interval_seconds. """ @@ -191,8 +211,8 @@ def decorator(poll_step_function: Callable[..., Any]) -> Callable[..., Any]: async def wrapped(*args: Any, **kwargs: Any) -> None: poll_interval_seconds = float(kwargs.pop(poll_interval_kwarg_name, default_poll_interval_seconds)) - async def loop_body() -> bool: - result = poll_step_function(*args, **kwargs) + async def loop_body(abort_signal: Any) -> bool: + result = poll_step_function(*args, abort_signal=abort_signal, **kwargs) should_continue = await result if inspect.isawaitable(result) else result return bool(should_continue) @@ -200,7 +220,6 @@ async def loop_body() -> bool: loop_body, poll_interval_seconds, channel_name=channel_name, - check_interval_seconds=check_interval_seconds, show_controls=show_controls, ) diff --git a/src/py/mat3ra/notebooks_utils/token_store.py b/src/py/mat3ra/notebooks_utils/token_store.py index 73eebd6a8..79a135d5c 100644 --- a/src/py/mat3ra/notebooks_utils/token_store.py +++ b/src/py/mat3ra/notebooks_utils/token_store.py @@ -31,3 +31,9 @@ async def load_token(oidc_url: str) -> Optional[dict]: if not entry or entry.get("expires_at", 0) <= time.time() + _EXPIRY_BUFFER: return None return entry + + +async def delete_token(oidc_url: str) -> None: + cache = await _read() + cache.pop(oidc_url, None) + await _write(cache) diff --git a/tests/py/unit/core/entity/test_job_api.py b/tests/py/unit/core/entity/test_job_api.py index d676a45b5..fc3c575e9 100644 --- a/tests/py/unit/core/entity/test_job_api.py +++ b/tests/py/unit/core/entity/test_job_api.py @@ -1,11 +1,14 @@ +import asyncio from typing import Any, Dict, List from unittest.mock import MagicMock import pytest +from mat3ra.api_client import AuthContext, JobEndpoints from mat3ra.notebooks_utils.core.entity.job.api import ( create_job, find_job_for_material, find_job_for_material_with_property, + get_jobs_statuses_by_ids_async, get_kgrid_of_job, get_kgrid_query, ) @@ -251,3 +254,27 @@ def test_find_job_for_material_matches_the_kgrid_of_the_unit(): ) def test_get_kgrid_of_job(unit, expected_kgrid): assert get_kgrid_of_job({"workflow": {"subworkflows": [{"units": [unit]}]}}) == expected_kgrid + + +REQUEST_TIMEOUT_SECONDS = 0.05 +BLOCKED_REQUEST_SECONDS = 1.0 +ACCESS_TOKEN = "access-token-1" +ABORT_SIGNAL = "abort-signal" +JOBS_FETCH_URL = ( + "https://platform.mat3ra.com:443/api/2018-10-01/jobs" + "?query=%7B%22_id%22%3A+%7B%22%24in%22%3A+%5B%22job-1%22%5D%7D%7D" + "&projection=%7B%22fields%22%3A+%7B%22status%22%3A+1%7D%7D" +) +BEARER_HEADERS = {"Authorization": f"Bearer {ACCESS_TOKEN}", "Content-Type": "application/json"} + + +@pytest.mark.asyncio +async def test_get_jobs_statuses_by_ids_async(monkeypatch): + pyfetch = MagicMock(side_effect=lambda *args, **kwargs: asyncio.sleep(BLOCKED_REQUEST_SECONDS)) + monkeypatch.setattr("mat3ra.notebooks_utils.core.entity.job.api.pyfetch", pyfetch) + monkeypatch.setattr("mat3ra.notebooks_utils.core.entity.job.api.is_pyodide_environment", lambda: True) + endpoint = JobEndpoints("platform.mat3ra.com", 443, OWNER_ID, None, auth=AuthContext(access_token=ACCESS_TOKEN)) + + with pytest.raises(asyncio.TimeoutError): + await get_jobs_statuses_by_ids_async(endpoint, [CREATED_JOB["_id"]], REQUEST_TIMEOUT_SECONDS, ABORT_SIGNAL) + pyfetch.assert_called_once_with(JOBS_FETCH_URL, headers=BEARER_HEADERS, signal=ABORT_SIGNAL) diff --git a/tests/py/unit/test_auth_retry.py b/tests/py/unit/test_auth_retry.py new file mode 100644 index 000000000..cd45fcf0d --- /dev/null +++ b/tests/py/unit/test_auth_retry.py @@ -0,0 +1,71 @@ +import asyncio +import contextlib +import os +from unittest.mock import MagicMock + +import pytest +import requests +from mat3ra.api_client import ACCESS_TOKEN_ENV_VAR, AuthContext +from mat3ra.notebooks_utils import auth, token_store +from mat3ra.notebooks_utils.api.job import wait_for_jobs_to_finish_async +from mat3ra.notebooks_utils.core.api import token_store as file_token_store + +POLL_INTERVAL_SECONDS = 0.01 +CACHED_TOKEN_DATA = {"access_token": "cached-token", "expires_in": 3600} +NEW_TOKEN_DATA = {"access_token": "new-token", "expires_in": 3600} +FINISHED_JOBS = [{"status": "finished"}] +HTTP_ERROR_401 = requests.HTTPError(response=MagicMock(status_code=401)) +HTTP_ERROR_403 = requests.HTTPError(response=MagicMock(status_code=403)) +HTTP_ERROR_503 = requests.HTTPError(response=MagicMock(status_code=503)) + + +async def completed_device_login(show_popup): + auth.store_token_data_in_environment(NEW_TOKEN_DATA) + return NEW_TOKEN_DATA + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status_results", "expectation", "expected_access_token"), + [ + ([asyncio.TimeoutError(), FINISHED_JOBS], contextlib.nullcontext(), "cached-token"), + ([HTTP_ERROR_503, FINISHED_JOBS], contextlib.nullcontext(), "cached-token"), + ([HTTP_ERROR_403, FINISHED_JOBS], pytest.raises(requests.HTTPError), "cached-token"), + ([HTTP_ERROR_401, FINISHED_JOBS], contextlib.nullcontext(), "new-token"), + ([HTTP_ERROR_401, HTTP_ERROR_401, FINISHED_JOBS], pytest.raises(requests.HTTPError), "new-token"), + ], +) +async def test_wait_for_jobs_to_finish_async(monkeypatch, tmp_path, status_results, expectation, expected_access_token): + monkeypatch.setattr(file_token_store, "_FILE_PATH", str(tmp_path / "oidc_token_cache.json")) + monkeypatch.setattr(auth, "authenticate_oidc", completed_device_login) + monkeypatch.setenv(ACCESS_TOKEN_ENV_VAR, CACHED_TOKEN_DATA["access_token"]) + auth_context = AuthContext(access_token=CACHED_TOKEN_DATA["access_token"]) + endpoint = MagicMock(auth=auth_context, list=MagicMock(side_effect=status_results)) + + with expectation: + await wait_for_jobs_to_finish_async(endpoint, ["job-1"], poll_interval=POLL_INTERVAL_SECONDS) + assert auth_context.access_token == expected_access_token + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("environment_token", "token_check_error", "checked_token", "expected_token"), + [ + ("", HTTP_ERROR_401, "cached-token", "new-token"), + ("environment-token", KeyError("data"), "environment-token", "new-token"), + ("", None, "cached-token", "cached-token"), + ], +) +async def test_authenticate(monkeypatch, tmp_path, environment_token, token_check_error, checked_token, expected_token): + monkeypatch.setattr(file_token_store, "_FILE_PATH", str(tmp_path / "oidc_token_cache.json")) + monkeypatch.setattr(auth, "authenticate_oidc", completed_device_login) + api_client = MagicMock() + api_client.authenticate.return_value.list_accounts.side_effect = token_check_error + monkeypatch.setattr(auth, "APIClient", api_client) + monkeypatch.setenv(ACCESS_TOKEN_ENV_VAR, environment_token) + await token_store.save_token(auth.get_oidc_base_url(), CACHED_TOKEN_DATA) + + await auth.authenticate(globals_dict={}) + + assert os.environ[ACCESS_TOKEN_ENV_VAR] == expected_token + api_client.authenticate.assert_called_once_with(access_token=checked_token) diff --git a/tests/py/unit/test_jupyterlite_interrupts.py b/tests/py/unit/test_jupyterlite_interrupts.py index 9b072e2b1..8f56ef64c 100644 --- a/tests/py/unit/test_jupyterlite_interrupts.py +++ b/tests/py/unit/test_jupyterlite_interrupts.py @@ -1,6 +1,11 @@ import asyncio +import contextlib +import threading +from unittest.mock import MagicMock import pytest +from mat3ra.notebooks_utils.api.job import wait_for_jobs_to_finish_async +from mat3ra.notebooks_utils.core.api.auth import _poll_for_token_data from mat3ra.notebooks_utils.pyodide.runtime import ( UserAbortError, interruptible_polling_loop, @@ -8,14 +13,21 @@ ) POLL_INTERVAL_SECONDS = 0.01 -CHECK_INTERVAL_SECONDS = 0.005 +ABORT_AFTER_SECONDS = 0.05 +BLOCKED_REQUEST_SECONDS = 1.0 +TOKEN_DATA = {"access_token": "new-token", "expires_in": 3600} +PENDING_TOKEN_RESPONSE = MagicMock(status_code=400, json=MagicMock(return_value={"error": "authorization_pending"})) +AUTHORIZED_TOKEN_RESPONSE = MagicMock(status_code=200, json=MagicMock(return_value=TOKEN_DATA)) +SERVER_ERROR_RESPONSE = MagicMock(status_code=502, json=MagicMock(side_effect=ValueError)) +REFUSED_TOKEN_RESPONSE = MagicMock(status_code=400, json=MagicMock(return_value={"error": "access_denied"})) +DEVICE_FLOW_ARGUMENTS = ("https://platform.mat3ra.com/oidc", "client-1", "device-code-1", POLL_INTERVAL_SECONDS, 600) @pytest.mark.asyncio async def test_run_interruptible_loop_async_stops_when_body_returns_false(): call_count = 0 - async def loop_body(): + async def loop_body(abort_signal): nonlocal call_count call_count += 1 return call_count < 3 @@ -24,17 +36,53 @@ async def loop_body(): loop_body, POLL_INTERVAL_SECONDS, show_controls=False, - check_interval_seconds=CHECK_INTERVAL_SECONDS, ) assert call_count == 3 +@pytest.mark.asyncio +@pytest.mark.parametrize( + "start_polling", + [ + lambda blocked_request: wait_for_jobs_to_finish_async(MagicMock(list=blocked_request), ["job-1"]), + lambda blocked_request: _poll_for_token_data(*DEVICE_FLOW_ARGUMENTS), + ], +) +async def test_abort_raises_user_abort_error_while_a_request_is_in_flight(monkeypatch, start_polling): + release_request = threading.Event() + + def blocked_request(*args, **kwargs): + release_request.wait(BLOCKED_REQUEST_SECONDS) + + monkeypatch.setattr("requests.post", blocked_request) + task = asyncio.create_task(start_polling(blocked_request)) + await asyncio.sleep(ABORT_AFTER_SECONDS) + task.cancel() + with pytest.raises(UserAbortError): + await task + release_request.set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("token_responses", "expectation"), + [ + ([PENDING_TOKEN_RESPONSE, SERVER_ERROR_RESPONSE, AUTHORIZED_TOKEN_RESPONSE], contextlib.nullcontext()), + ([REFUSED_TOKEN_RESPONSE], pytest.raises(Exception, match=r"\(400\): access_denied")), + ], +) +async def test_poll_for_token_data(monkeypatch, token_responses, expectation): + monkeypatch.setattr("requests.post", MagicMock(side_effect=token_responses)) + with expectation: + assert await _poll_for_token_data(*DEVICE_FLOW_ARGUMENTS) == TOKEN_DATA + + @pytest.mark.asyncio async def test_interruptible_polling_loop_decorator_returns_coroutine_and_runs_until_false(): call_count = 0 @interruptible_polling_loop(show_controls=False) - def poll_step(): + def poll_step(abort_signal): nonlocal call_count call_count += 1 return call_count < 2