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