diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 6d63ccda..b8034737 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -333,6 +333,10 @@ jobs: - name: Install project run: poetry install --no-interaction --only-root + - name: Install Streamlit for widget integration tests + if: matrix.python-version == '3.12' + run: poetry install --no-interaction --extras server + - name: Run unit tests env: TOOLKIT_VERSION: ${{ steps.version.outputs.VERSION }} diff --git a/README.md b/README.md index 5a8cab41..b911a97b 100644 --- a/README.md +++ b/README.md @@ -22,6 +22,8 @@ It starts and manages Jupyter, Streamlit, and LSP servers, and provides runtime - Native **Deepnote component library** including beautiful `DataFrame` rendering and interactive inputs - **Python kernel with curated set of libraries preinstalled**, allowing you to focus on work instead of fighting with Python dependencies - Run multiple **interactive applications built with Streamlit** +- Build custom Streamlit interfaces over local Deepnote files and hosted runs with + [per-viewer authentication](docs/streamlit-apps.md) - Language Server Protocol integration for code completion and intelligence - Git integration with SSH/HTTPS authentication diff --git a/deepnote_toolkit/notebooks/__init__.py b/deepnote_toolkit/notebooks/__init__.py new file mode 100644 index 00000000..af7de61f --- /dev/null +++ b/deepnote_toolkit/notebooks/__init__.py @@ -0,0 +1,28 @@ +"""Read `.deepnote` files and run notebooks, independent of any UI framework. + +Only the names in __all__ are supported public API. Wire schemas and API clients +are implementation details. +""" + +from .cloud_runner import DeepnoteCloudRunner +from .credentials import ApiCredentials, CredentialsProvider +from .document import DeepnoteDocument +from .local_runner import DeepnoteLocalRunner +from .models import DeepnoteDataframe, InputBlock, NotebookOutput, RunnerInfo +from .run_result import RunResult +from .runner import Runner, RunnerError + +__all__ = [ + "ApiCredentials", + "CredentialsProvider", + "DeepnoteCloudRunner", + "DeepnoteDataframe", + "DeepnoteDocument", + "DeepnoteLocalRunner", + "InputBlock", + "NotebookOutput", + "RunResult", + "Runner", + "RunnerError", + "RunnerInfo", +] diff --git a/deepnote_toolkit/notebooks/_schemas.py b/deepnote_toolkit/notebooks/_schemas.py new file mode 100644 index 00000000..f9833bb5 --- /dev/null +++ b/deepnote_toolkit/notebooks/_schemas.py @@ -0,0 +1,62 @@ +"""Consumed fields of the v2 API contracts (contracts/runs.ts and notebooks.ts). + +Extra fields are ignored. +""" + +from __future__ import annotations + +from typing import Any + +from pydantic import BaseModel, Field, StrictBool, StrictFloat, StrictInt, StrictStr + +from .api_types import RunStatus, SnapshotStatus + + +class ApiInput(BaseModel): + name: StrictStr + type: StrictStr + value: StrictStr | StrictBool | list[StrictStr] | None = None + label: StrictStr | None = None + options: list[StrictStr] = Field(default_factory=list) + multiple: StrictBool = False + min: StrictInt | StrictFloat | None = None + max: StrictInt | StrictFloat | None = None + step: StrictInt | StrictFloat | None = None + + +class ApiNotebook(BaseModel): + name: StrictStr = "Untitled notebook" + inputs: list[ApiInput] = Field(default_factory=list) + + +class NotebookResponse(BaseModel): + notebook: ApiNotebook + + +class CreateRunResponse(BaseModel): + """The run identity and execution status returned by POST /v2/runs.""" + + run_id: StrictStr = Field(alias="runId", min_length=1) + status: RunStatus + + +class ApiRun(CreateRunResponse): + """GET run details, including the required snapshot lifecycle status.""" + + snapshot_status: SnapshotStatus = Field(alias="snapshotStatus") + snapshot_blocks: list[dict[str, Any]] | None = Field( + default=None, alias="snapshotBlocks" + ) + error: StrictStr | None = None + + +class GetRunResponse(BaseModel): + """The response envelope returned when fetching an existing run.""" + + run: ApiRun + + +class ViewerTokenResponse(BaseModel): + token: StrictStr = Field(min_length=1) + api_origin: StrictStr = Field(alias="apiOrigin") + expires_at_seconds: StrictInt | StrictFloat = Field(alias="expiresAtSeconds") diff --git a/deepnote_toolkit/notebooks/api_client.py b/deepnote_toolkit/notebooks/api_client.py new file mode 100644 index 00000000..43494b67 --- /dev/null +++ b/deepnote_toolkit/notebooks/api_client.py @@ -0,0 +1,221 @@ +"""The Deepnote public API operations a runner needs.""" + +from __future__ import annotations + +import time +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from typing import Any, TypeVar, cast +from urllib.parse import quote + +import requests +from pydantic import BaseModel, ValidationError + +from ._schemas import ( + ApiInput, + ApiRun, + CreateRunResponse, + GetRunResponse, + NotebookResponse, +) +from .api_types import ( + INPUT_BLOCK_TYPES, + TERMINAL_RUN_STATUSES, + InputBlockType, + InputValue, + RunStatus, + SnapshotStatus, + StorageMode, +) +from .credentials import CredentialsProvider +from .models import InputBlock, NotebookOutput +from .runner import RunnerError +from .transport import request_json +from .wire import decode_block_outputs + +Schema = TypeVar("Schema", bound=BaseModel) + + +@dataclass(frozen=True) +class CloudNotebook: + """A notebook's name and input blocks.""" + + name: str + inputs: tuple[InputBlock, ...] + + +@dataclass(frozen=True) +class CloudRun: + """The state of one run. `outputs` is None until the run's snapshot is stored.""" + + run_id: str + status: RunStatus + snapshot_status: SnapshotStatus | None + outputs: tuple[NotebookOutput, ...] | None + error: str | None + + @property + def is_finished(self) -> bool: + """Whether the run has reached a final status.""" + + return self.status in TERMINAL_RUN_STATUSES + + +class DeepnoteApiClient: + """Sends requests to the Deepnote public API and validates what comes back.""" + + def __init__( + self, + credentials: CredentialsProvider, + *, + session: requests.Session | None = None, + request_timeout: float = 30, + clock: Callable[[], float] = time.monotonic, + ): + self._credentials = credentials + self._clock = clock + self._session = session if session is not None else requests.Session() + self._request_timeout = request_timeout + + def get_notebook(self, notebook_id: str) -> CloudNotebook: + """Read a notebook's name and input blocks.""" + + payload = self._request("GET", f"/v2/notebooks/{quote(notebook_id, safe='')}") + notebook = _validate(NotebookResponse, payload, "notebook").notebook + return CloudNotebook( + name=notebook.name, + inputs=tuple( + _input_block(value) + for value in notebook.inputs + if value.type in INPUT_BLOCK_TYPES and value.name + ), + ) + + def create_run( + self, + notebook_id: str, + inputs: Mapping[str, Any], + *, + storage_mode: StorageMode | None = None, + timeout: float | None = None, + ) -> CloudRun: + """Start a detached run of the whole notebook.""" + + body: dict[str, Any] = { + "notebookId": notebook_id, + "detached": True, + "inputs": { + name: _encode_input(name, value) for name, value in inputs.items() + }, + } + if storage_mode is not None: + body["detachedRunStorageMode"] = storage_mode + payload = self._request("POST", "/v2/runs", body, timeout=timeout) + run = _validate(CreateRunResponse, payload, "run") + return CloudRun( + run_id=run.run_id, + status=run.status, + snapshot_status=None, + outputs=None, + error=None, + ) + + def get_run(self, run_id: str, *, timeout: float | None = None) -> CloudRun: + """Read a run with the outputs of the notebook it executed.""" + + # The blocks delivery holds the executed notebook alone, not the whole project. + payload = self._request( + "GET", + f"/v2/runs/{quote(run_id, safe='')}?snapshotDelivery=blocks", + timeout=timeout, + ) + return _cloud_run(_validate(GetRunResponse, payload, "run").run) + + def _request( + self, + method: str, + path: str, + body: Mapping[str, Any] | None = None, + *, + timeout: float | None = None, + ) -> Mapping[str, Any]: + budget = ( + self._request_timeout + if timeout is None + else min(timeout, self._request_timeout) + ) + if budget <= 0: + raise RunnerError("API request deadline expired", transient=True) + deadline = self._clock() + budget + credentials = self._credentials(timeout=budget) + remaining = deadline - self._clock() + if remaining <= 0: + raise RunnerError( + "API credentials exhausted the request timeout", transient=True + ) + return request_json( + self._session, + method, + f"{credentials.api_origin}{path}", + headers={"Authorization": f"Bearer {credentials.token}"}, + body=body, + timeout=remaining, + ) + + +def _validate(schema: type[Schema], payload: Mapping[str, Any], what: str) -> Schema: + """Validate an API response, reporting schema failures as `RunnerError`.""" + + try: + return schema(**payload) + except ValidationError as error: + raise RunnerError( + f"Deepnote API returned an invalid {what} response" + ) from error + + +def _encode_input(name: str, value: Any) -> InputValue: + """Convert a value to the form the runs API accepts, or raise `ValueError`.""" + + if isinstance(value, bool): + return value + if isinstance(value, (list, tuple)): + return [str(item) for item in value] + if value is None or isinstance(value, (Mapping, set, frozenset, bytes)): + raise ValueError( + f'Input "{name}" has a {type(value).__name__} value. ' + "Pass text, a number, a boolean or a list of texts." + ) + return str(value) + + +def _input_block(value: ApiInput) -> InputBlock: + """Convert validated API input metadata to a notebook input block.""" + + return InputBlock( + variable_name=value.name, + type=cast(InputBlockType, value.type), + value=value.value, + label=value.label, + options=tuple(value.options), + multiple=value.multiple, + min=value.min, + max=value.max, + step=value.step, + ) + + +def _cloud_run(run: ApiRun) -> CloudRun: + """Convert a validated run and any available snapshot blocks to runner data.""" + + return CloudRun( + run_id=run.run_id, + status=run.status, + snapshot_status=run.snapshot_status, + outputs=( + decode_block_outputs(run.snapshot_blocks, id_key="id") + if run.snapshot_blocks is not None + else None + ), + error=run.error, + ) diff --git a/deepnote_toolkit/notebooks/api_types.py b/deepnote_toolkit/notebooks/api_types.py new file mode 100644 index 00000000..1a9623e4 --- /dev/null +++ b/deepnote_toolkit/notebooks/api_types.py @@ -0,0 +1,27 @@ +"""Value types of the Deepnote public API.""" + +from __future__ import annotations + +from typing import Literal, Union, get_args + +InputBlockType = Literal[ + "input-checkbox", + "input-date", + "input-date-range", + "input-file", + "input-select", + "input-slider", + "input-text", + "input-textarea", +] +RunStatus = Literal[ + "pending", "running", "success", "error", "internal_error", "stopped" +] +SnapshotStatus = Literal["pending", "available", "unavailable"] +StorageMode = Literal["read_write", "readonly"] +InputValue = Union[str, bool, list[str]] + +INPUT_BLOCK_TYPES: frozenset[InputBlockType] = frozenset(get_args(InputBlockType)) +TERMINAL_RUN_STATUSES: frozenset[RunStatus] = frozenset( + {"success", "error", "internal_error", "stopped"} +) diff --git a/deepnote_toolkit/notebooks/cloud_runner.py b/deepnote_toolkit/notebooks/cloud_runner.py new file mode 100644 index 00000000..95d9deb1 --- /dev/null +++ b/deepnote_toolkit/notebooks/cloud_runner.py @@ -0,0 +1,166 @@ +"""Run a notebook in Deepnote Cloud through the public API.""" + +from __future__ import annotations + +import math +import time +from collections.abc import Callable, Mapping +from typing import Any + +import requests + +from .api_client import CloudRun, DeepnoteApiClient +from .api_types import StorageMode +from .credentials import ( + DEFAULT_API_ORIGIN, + CredentialsProvider, + TokenProvider, + token_credentials, +) +from .models import RunnerInfo +from .run_result import RunResult +from .runner import RunnerError + +Sleep = Callable[[float], None] + +MAX_TRANSIENT_POLL_FAILURES = 5 + + +class DeepnoteCloudRunner: + """Run an existing notebook directly through the Deepnote public API. + + The token comes from `token`, `token_provider` or the `DEEPNOTE_TOKEN` + environment variable, and is sent to `base_url`. A token provider is called for + every request, which lets a long-lived process use short-lived credentials. + `credentials` replaces all three with one provider of the token and its origin. + + `storage_mode="readonly"` keeps the run from changing the project's stored + files. None leaves the choice to the API, which allows writes. + + `timeout` is a polling budget, including credentials and API requests. Network + operations or custom credential providers can overrun it. Expiry stops polling + but does not cancel the notebook. + + The outputs can arrive after the run finishes. `snapshot_timeout` is how many + seconds to wait for them. A result whose `snapshot_status` is still `pending` + has none because that wait ran out. + """ + + def __init__( + self, + notebook_id: str, + *, + token: str | None = None, + token_provider: TokenProvider | None = None, + base_url: str = DEFAULT_API_ORIGIN, + credentials: CredentialsProvider | None = None, + storage_mode: StorageMode | None = None, + timeout: float = 600, + snapshot_timeout: float = 10, + poll_interval: float = 2, + session: requests.Session | None = None, + sleep: Sleep = time.sleep, + clock: Callable[[], float] = time.monotonic, + ): + if not notebook_id: + raise ValueError("notebook_id is required") + if not math.isfinite(timeout) or timeout <= 0: + raise ValueError("timeout must be positive and finite") + if not math.isfinite(snapshot_timeout) or snapshot_timeout < 0: + raise ValueError("snapshot_timeout must be non-negative and finite") + if not math.isfinite(poll_interval) or poll_interval <= 0: + raise ValueError("poll_interval must be positive") + if credentials is not None and ( + token is not None or token_provider is not None + ): + raise ValueError("Pass credentials or a token, not both") + self.notebook_id = notebook_id + self.storage_mode = storage_mode + self.timeout = timeout + self.snapshot_timeout = snapshot_timeout + self.poll_interval = poll_interval + self._client = DeepnoteApiClient( + credentials or token_credentials(token, token_provider, base_url=base_url), + session=session, + request_timeout=min(timeout, 30), + clock=clock, + ) + self._sleep = sleep + self._clock = clock + + def info(self) -> RunnerInfo: + """Read the notebook's name and input blocks from the public API.""" + + notebook = self._client.get_notebook(self.notebook_id) + return RunnerInfo( + notebook=notebook.name, inputs=notebook.inputs, run_target="cloud" + ) + + def run(self, inputs: Mapping[str, Any]) -> RunResult: + """Start a detached run with the given input values and wait for its result.""" + + deadline = self._clock() + self.timeout + run = self._client.create_run( + self.notebook_id, + inputs, + storage_mode=self.storage_mode, + timeout=self.timeout, + ) + if self._clock() >= deadline: + self._timed_out(run) + run = self._wait_until_finished(run, deadline) + run = self._settle_snapshot( + run, min(deadline, self._clock() + self.snapshot_timeout) + ) + return RunResult( + target="cloud", + success=run.status == "success", + outputs=run.outputs or (), + run_id=run.run_id, + status=run.status, + error=run.error, + snapshot_status=run.snapshot_status, + ) + + def _timed_out(self, run: CloudRun) -> None: + raise RunnerError( + f"Deepnote run {run.run_id} did not finish in {self.timeout:g} seconds" + ) + + def _pause(self, deadline: float) -> bool: + remaining = deadline - self._clock() + if remaining <= 0: + return False + self._sleep(min(self.poll_interval, remaining)) + return self._clock() < deadline + + def _wait_until_finished(self, run: CloudRun, deadline: float) -> CloudRun: + transient_failures = 0 + while not run.is_finished: + if not self._pause(deadline): + self._timed_out(run) + try: + run = self._client.get_run(run.run_id, timeout=deadline - self._clock()) + transient_failures = 0 + except RunnerError as error: + transient_failures += 1 + if ( + not error.transient + or transient_failures > MAX_TRANSIENT_POLL_FAILURES + ): + raise + if self._clock() >= deadline: + self._timed_out(run) + return run + + def _settle_snapshot(self, run: CloudRun, deadline: float) -> CloudRun: + # POST has no snapshot metadata, even when it reports a completed run. + while run.snapshot_status in {None, "pending"}: + if not self._pause(deadline): + break + try: + run = self._client.get_run(run.run_id, timeout=deadline - self._clock()) + except RunnerError as error: + if not error.transient: + raise + return run diff --git a/deepnote_toolkit/notebooks/credentials.py b/deepnote_toolkit/notebooks/credentials.py new file mode 100644 index 00000000..3257bf7c --- /dev/null +++ b/deepnote_toolkit/notebooks/credentials.py @@ -0,0 +1,54 @@ +"""Where a runner gets its API token and the origin to send it to.""" + +from __future__ import annotations + +import os +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Protocol + +from .runner import RunnerError + +DEFAULT_API_ORIGIN = "https://api.deepnote.com" +TokenProvider = Callable[[], str] + + +@dataclass(frozen=True) +class ApiCredentials: + """A bearer token and the API origin it is valid at.""" + + token: str = field(repr=False) + api_origin: str = DEFAULT_API_ORIGIN + + +class CredentialsProvider(Protocol): + """Returns the credentials for one request. Called before every request.""" + + def __call__(self, *, timeout: float = 30) -> ApiCredentials: + """Return the credentials, or raise `RunnerError` when there are none.""" + + +def token_credentials( + token: str | None = None, + token_provider: TokenProvider | None = None, + *, + base_url: str = DEFAULT_API_ORIGIN, +) -> CredentialsProvider: + """Credentials from `token`, `token_provider` or `DEEPNOTE_TOKEN`, in that order.""" + + if token is not None and token_provider is not None: + raise ValueError("Pass token or token_provider, not both") + api_origin = base_url.rstrip("/") + + def provide(*, timeout: float = 30) -> ApiCredentials: + if token_provider is not None: + value: str | None = token_provider() + elif token is not None: + value = token + else: + value = os.environ.get("DEEPNOTE_TOKEN") + if not value: + raise RunnerError("A Deepnote API token is required") + return ApiCredentials(token=value, api_origin=api_origin) + + return provide diff --git a/deepnote_toolkit/notebooks/document.py b/deepnote_toolkit/notebooks/document.py new file mode 100644 index 00000000..e2417ccf --- /dev/null +++ b/deepnote_toolkit/notebooks/document.py @@ -0,0 +1,111 @@ +"""Read `.deepnote` source and snapshot files.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from pathlib import Path +from typing import Any + +import yaml + +from .api_types import INPUT_BLOCK_TYPES +from .models import InputBlock, NotebookOutput +from .outputs import OutputCollection +from .wire import decode_block_outputs, optional_number, optional_string, string_tuple +from .yaml_loader import load_yaml + + +class DeepnoteDocument(OutputCollection): + """A parsed source or snapshot `.deepnote` file. + + `notebook_id` limits the inputs and outputs to one notebook of the project. + """ + + def __init__(self, raw: Mapping[str, Any], *, notebook_id: str | None = None): + project = raw.get("project") + if not isinstance(project, Mapping) or not isinstance( + project.get("notebooks"), list + ): + raise ValueError("Expected a .deepnote document with project.notebooks") + notebooks = project["notebooks"] + if notebook_id is not None: + notebooks = [ + notebook + for notebook in notebooks + if isinstance(notebook, Mapping) and notebook.get("id") == notebook_id + ] + if not notebooks: + raise ValueError(f"Notebook {notebook_id} is not in this document") + self.raw = raw + self.project_name = str(project.get("name", "Untitled project")) + self.inputs, self.outputs = _read_blocks(notebooks) + + @classmethod + def load( + cls, path: str | Path, *, notebook_id: str | None = None + ) -> DeepnoteDocument: + """Read a `.deepnote` file from disk. Raises `ValueError` on invalid content.""" + + source = Path(path) + try: + return cls.parse( + source.read_text(encoding="utf-8"), notebook_id=notebook_id + ) + except ValueError as error: + raise ValueError(f"{source}: {error}") from None + + @classmethod + def parse(cls, content: str, *, notebook_id: str | None = None) -> DeepnoteDocument: + """Read `.deepnote` YAML from a string. Raises `ValueError` when invalid.""" + + try: + raw = load_yaml(content) + except yaml.YAMLError as error: + raise ValueError(f"Could not parse .deepnote YAML: {error}") from error + if not isinstance(raw, Mapping): + raise ValueError("Expected .deepnote YAML to contain an object") + return cls(raw, notebook_id=notebook_id) + + +def _read_blocks( + notebooks: Sequence[Any], +) -> tuple[tuple[InputBlock, ...], tuple[NotebookOutput, ...]]: + inputs: list[InputBlock] = [] + outputs: list[NotebookOutput] = [] + for notebook in notebooks: + if not isinstance(notebook, Mapping): + continue + blocks = notebook.get("blocks") + if not isinstance(blocks, list): + continue + inputs.extend( + input_block + for block in blocks + if isinstance(block, Mapping) and (input_block := _read_input_block(block)) + ) + outputs.extend(decode_block_outputs(blocks, id_key="id")) + return tuple(inputs), tuple(outputs) + + +def _read_input_block(block: Mapping[str, Any]) -> InputBlock | None: + block_type = str(block.get("type", "")) + metadata = block.get("metadata") + if block_type not in INPUT_BLOCK_TYPES or not isinstance(metadata, Mapping): + return None + variable_name = metadata.get("deepnote_variable_name") + if not isinstance(variable_name, str) or not variable_name: + return None + return InputBlock( + variable_name=variable_name, + type=block_type, + label=optional_string(metadata.get("deepnote_input_label")), + value=metadata.get("deepnote_variable_value"), + options=string_tuple(metadata.get("deepnote_variable_options")), + options_from_variable=( + metadata.get("deepnote_variable_select_type") == "from-variable" + ), + multiple=metadata.get("deepnote_allow_multiple_values") is True, + min=optional_number(metadata.get("deepnote_slider_min_value")), + max=optional_number(metadata.get("deepnote_slider_max_value")), + step=optional_number(metadata.get("deepnote_slider_step")), + ) diff --git a/deepnote_toolkit/notebooks/local_runner.py b/deepnote_toolkit/notebooks/local_runner.py new file mode 100644 index 00000000..db2d6a80 --- /dev/null +++ b/deepnote_toolkit/notebooks/local_runner.py @@ -0,0 +1,90 @@ +"""Run a notebook through a local `@deepnote/local-runner` sidecar.""" + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import Any + +import requests + +from .document import DeepnoteDocument +from .models import RunnerInfo +from .run_result import RunResult +from .transport import request_json +from .wire import decode_block_outputs, decode_inputs, optional_string + +logger = logging.getLogger(__name__) + + +class DeepnoteLocalRunner: + """Run a notebook through a local sidecar, configured for a local or cloud kernel.""" + + def __init__( + self, + base_url: str = "http://127.0.0.1:8787", + *, + timeout: float = 600, + session: requests.Session | None = None, + ): + self.base_url = base_url.rstrip("/") + self.timeout = timeout + self._session = session if session is not None else requests.Session() + + def info(self) -> RunnerInfo: + """Read the notebook's name and input blocks from the sidecar.""" + + payload = self._request("GET", "/api/info") + return RunnerInfo( + notebook=str(payload.get("notebook", "Untitled project")), + inputs=decode_inputs(payload.get("inputs")), + run_target=str(payload.get("runTarget", "")), + ) + + def run(self, inputs: Mapping[str, Any]) -> RunResult: + """Run the notebook in the sidecar with the given input values.""" + + return _decode_run_result( + self._request("POST", "/api/run", {"inputs": dict(inputs)}) + ) + + def _request( + self, method: str, path: str, body: Mapping[str, Any] | None = None + ) -> Mapping[str, Any]: + return request_json( + self._session, + method, + f"{self.base_url}{path}", + headers={}, + body=body, + timeout=self.timeout, + ) + + +def _decode_run_result(payload: Mapping[str, Any]) -> RunResult: + snapshot = None + snapshot_yaml = payload.get("snapshotYaml") + if isinstance(snapshot_yaml, str) and snapshot_yaml: + try: + snapshot = DeepnoteDocument.parse(snapshot_yaml) + except ValueError: + logger.warning( + "Could not parse the run snapshot; using inline outputs (runId=%r, target=%r)", + optional_string(payload.get("runId")), + optional_string(payload.get("target")), + ) + return RunResult( + target=str(payload.get("target", "")), + success=payload.get("success") is True, + outputs=( + snapshot.outputs + if snapshot + else decode_block_outputs(payload.get("outputs"), id_key="blockId") + ), + run_id=optional_string(payload.get("runId")), + status=optional_string(payload.get("status")), + error=optional_string(payload.get("error")), + view_url=optional_string(payload.get("viewUrl")), + snapshot=snapshot, + created=payload.get("created") is True, + ) diff --git a/deepnote_toolkit/notebooks/models.py b/deepnote_toolkit/notebooks/models.py new file mode 100644 index 00000000..b53ea26a --- /dev/null +++ b/deepnote_toolkit/notebooks/models.py @@ -0,0 +1,195 @@ +"""Typed models for Deepnote input blocks, outputs and runner metadata.""" + +from __future__ import annotations + +import base64 +import binascii +from collections.abc import Iterable, Mapping +from dataclasses import dataclass +from typing import Any + +from .api_types import InputBlockType + +DATAFRAME_MIME = "application/vnd.deepnote.dataframe.v3+json" +INDEX_COLUMN = "_deepnote_index_column" + + +def join_text(value: Any) -> str: + """Normalize nbformat's string-or-list text values to one string.""" + + if isinstance(value, list): + return "".join(str(part) for part in value) + return "" if value is None else str(value) + + +@dataclass(frozen=True) +class InputBlock: + """The metadata a UI needs to render one Deepnote input block.""" + + variable_name: str + type: InputBlockType + value: Any + label: str | None = None + options: tuple[str, ...] = () + multiple: bool = False + min: float | int | None = None + max: float | int | None = None + step: float | int | None = None + options_from_variable: bool = False + + +@dataclass(frozen=True) +class DeepnoteDataframe: + """A structured Deepnote dataframe output, independent of pandas.""" + + columns: tuple[Mapping[str, Any], ...] + rows: tuple[Mapping[str, Any], ...] + raw: Mapping[str, Any] + + @classmethod + def from_value(cls, value: Any) -> DeepnoteDataframe | None: + """Read a dataframe output payload. Returns None when the value is not one.""" + + if not isinstance(value, Mapping): + return None + columns = value.get("columns") + rows = value.get("rows") + if not isinstance(columns, list) or not isinstance(rows, list): + return None + if not all(isinstance(column, Mapping) for column in columns): + return None + if not all(isinstance(row, Mapping) for row in rows): + return None + return cls(columns=tuple(columns), rows=tuple(rows), raw=value) + + @property + def row_count(self) -> int: + """Rows in the full dataframe. `rows` holds only the first page of them.""" + + count = self.raw.get("row_count") + return count if isinstance(count, int) else len(self.rows) + + @property + def is_truncated(self) -> bool: + """Whether the full dataframe has more rows than `rows` holds.""" + + return self.row_count > len(self.rows) + + def records(self, *, include_index: bool = True) -> list[dict[str, Any]]: + """Return rows as plain dicts, optionally without the index column.""" + + if include_index: + return [dict(row) for row in self.rows] + return [ + {key: value for key, value in row.items() if key != INDEX_COLUMN} + for row in self.rows + ] + + +@dataclass(frozen=True) +class NotebookOutput: + """One nbformat-compatible output emitted by a Deepnote block.""" + + block_id: str + block_type: str | None + raw: Mapping[str, Any] + + @property + def output_type(self) -> str: + """The nbformat output type, such as `stream` or `execute_result`.""" + + return str(self.raw.get("output_type", "")) + + @property + def data(self) -> Mapping[str, Any]: + """The output's MIME bundle, empty for outputs that have none.""" + + value = self.raw.get("data") + return value if isinstance(value, Mapping) else {} + + def text(self, mime: str = "text/plain") -> str: + """The output's text for a MIME type. Stream outputs count as `text/plain`.""" + + if self.output_type == "stream" and mime == "text/plain": + return join_text(self.raw.get("text")) + return join_text(self.data.get(mime)) + + def image_bytes(self, mime: str = "image/png") -> bytes | None: + """The decoded image for a MIME type, or None when absent or not base64.""" + + value = self.data.get(mime) + if value is None: + return None + encoded = "".join(join_text(value).split()) + try: + return base64.b64decode(encoded, validate=True) + except (ValueError, binascii.Error): + return None + + @property + def dataframe(self) -> DeepnoteDataframe | None: + """The output as a Deepnote dataframe, or None when it is not one.""" + + return DeepnoteDataframe.from_value(self.data.get(DATAFRAME_MIME)) + + +@dataclass(frozen=True) +class RunnerInfo: + """The target and input contract exposed by a Deepnote runner.""" + + notebook: str + inputs: tuple[InputBlock, ...] + run_target: str + + def matches_inputs(self, inputs: Iterable[InputBlock]) -> bool: + """Return whether the static input definitions match this runner's notebook. + + Names, block types, single or multiple selection, slider bounds and select + options must match. Options filled from a variable change between runs, so + they are not compared. + """ + + expected = tuple(inputs) + for blocks in (expected, self.inputs): + names = [block.variable_name for block in blocks] + if any(not name for name in names) or len(names) != len(set(names)): + return False + dynamic = frozenset( + input_block.variable_name + for input_block in expected + if input_block.options_from_variable + ) + return _input_contract(expected, dynamic) == _input_contract( + self.inputs, dynamic + ) + + +def _input_contract( + inputs: Iterable[InputBlock], dynamic_options: frozenset[str] +) -> frozenset[tuple[Any, ...]]: + return frozenset( + ( + input_block.variable_name, + input_block.type, + *_value_constraints(input_block, dynamic_options), + ) + for input_block in inputs + ) + + +def _value_constraints( + input_block: InputBlock, dynamic_options: frozenset[str] +) -> tuple[Any, ...]: + if input_block.type == "input-slider": + return ( + input_block.min if input_block.min is not None else 0, + input_block.max if input_block.max is not None else 100, + input_block.step if input_block.step is not None else 1, + ) + if input_block.type == "input-select": + is_dynamic = input_block.variable_name in dynamic_options + return ( + input_block.multiple, + None if is_dynamic else frozenset(input_block.options), + ) + return () diff --git a/deepnote_toolkit/notebooks/outputs.py b/deepnote_toolkit/notebooks/outputs.py new file mode 100644 index 00000000..252d4d0f --- /dev/null +++ b/deepnote_toolkit/notebooks/outputs.py @@ -0,0 +1,42 @@ +"""Queries shared by every source of notebook outputs.""" + +from __future__ import annotations + +from .models import DeepnoteDataframe, NotebookOutput + + +class OutputCollection: + """Shared output queries for a loaded document and a live run result.""" + + outputs: tuple[NotebookOutput, ...] + + def first_dataframe(self) -> DeepnoteDataframe | None: + """The first dataframe output, or None when there is none.""" + + for output in self.outputs: + if dataframe := output.dataframe: + return dataframe + return None + + def images(self, mime: str = "image/png") -> list[bytes]: + """Every image of the given MIME type, decoded.""" + + return [ + image + for output in self.outputs + if (image := output.image_bytes(mime)) is not None + ] + + def text(self, mime: str = "text/plain") -> str: + """The text of all outputs for a MIME type, joined.""" + + return "".join(output.text(mime) for output in self.outputs).strip() + + def agent_text(self) -> str: + """The text written by agent blocks, preferring Markdown over plain text.""" + + return "".join( + output.text("text/markdown") or output.text() + for output in self.outputs + if output.block_type == "agent" + ).strip() diff --git a/deepnote_toolkit/notebooks/run_result.py b/deepnote_toolkit/notebooks/run_result.py new file mode 100644 index 00000000..22af4fb4 --- /dev/null +++ b/deepnote_toolkit/notebooks/run_result.py @@ -0,0 +1,30 @@ +"""The result of one notebook run, from either runner.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from .api_types import SnapshotStatus +from .document import DeepnoteDocument +from .models import NotebookOutput +from .outputs import OutputCollection + + +@dataclass(frozen=True) +class RunResult(OutputCollection): + """What one run produced, whether it ran in Deepnote Cloud or locally. + + `snapshot_status` is set for cloud runs. `view_url`, `snapshot` and `created` + are set by the local runner. + """ + + target: str + success: bool + outputs: tuple[NotebookOutput, ...] = () + run_id: str | None = None + status: str | None = None + error: str | None = None + snapshot_status: SnapshotStatus | None = None + view_url: str | None = None + snapshot: DeepnoteDocument | None = None + created: bool = False diff --git a/deepnote_toolkit/notebooks/runner.py b/deepnote_toolkit/notebooks/runner.py new file mode 100644 index 00000000..6ceb1407 --- /dev/null +++ b/deepnote_toolkit/notebooks/runner.py @@ -0,0 +1,27 @@ +"""The interface every notebook runner implements.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, Protocol + +from .models import RunnerInfo +from .run_result import RunResult + + +class RunnerError(RuntimeError): + """The Deepnote runner was unavailable or rejected a request.""" + + def __init__(self, message: str, *, transient: bool = False): + super().__init__(message) + self.transient = transient + + +class Runner(Protocol): + """Runs one notebook and reports the inputs it accepts.""" + + def info(self) -> RunnerInfo: + """Return the notebook's name and the inputs it accepts.""" + + def run(self, inputs: Mapping[str, Any]) -> RunResult: + """Run the notebook with the given input values and wait for the result.""" diff --git a/deepnote_toolkit/notebooks/transport.py b/deepnote_toolkit/notebooks/transport.py new file mode 100644 index 00000000..bd5d1fdd --- /dev/null +++ b/deepnote_toolkit/notebooks/transport.py @@ -0,0 +1,83 @@ +"""Internal HTTP helpers shared by notebook execution and viewer authentication.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any +from urllib.parse import urlsplit + +import requests +from urllib3.util import Timeout + +from .runner import RunnerError + + +def _preserve_authorization( + request: requests.PreparedRequest, +) -> requests.PreparedRequest: + """Keep resolved credentials instead of applying Session.auth or .netrc.""" + return request + + +def request_json( + session: requests.Session, + method: str, + url: str, + *, + headers: Mapping[str, str], + body: Mapping[str, Any] | None = None, + timeout: float, +) -> Mapping[str, Any]: + """Send one request, without replaying POSTs or forwarding credentials on redirects.""" + origin = urlsplit(url) + origin_name = f"{origin.scheme}://{origin.netloc}" + request_headers = requests.structures.CaseInsensitiveDict( + {"Accept": "application/json", **headers} + ) + try: + with session.request( + method, + url, + headers=request_headers, + auth=( + _preserve_authorization if "Authorization" in request_headers else None + ), + json=body, + timeout=Timeout(total=timeout), + allow_redirects=False, + ) as response: + if 300 <= response.status_code < 400: + raise RunnerError( + f"{origin_name} returned HTTP {response.status_code}: " + "Refused a redirect" + ) + if response.status_code >= 400: + message = response.reason + try: + payload = response.json() + if isinstance(payload, Mapping): + reason = payload.get("message") or payload.get("error") + if isinstance(reason, str): + message = reason + except ValueError: + pass + raise RunnerError( + f"{origin_name} returned HTTP {response.status_code}: {message}", + transient=response.status_code == 429 + or response.status_code >= 500, + ) + try: + payload = response.json() + except ValueError as error: + raise RunnerError(f"{origin_name} returned invalid JSON") from error + except requests.Timeout as error: + raise RunnerError( + f"{origin_name} timed out after {timeout:g} seconds", transient=True + ) from error + except requests.RequestException as error: + raise RunnerError( + f"Could not reach {origin_name}: {error}", transient=True + ) from error + if not isinstance(payload, Mapping): + raise RunnerError(f"{origin_name} returned a non-object response") + return payload diff --git a/deepnote_toolkit/notebooks/wire.py b/deepnote_toolkit/notebooks/wire.py new file mode 100644 index 00000000..a9ef3c03 --- /dev/null +++ b/deepnote_toolkit/notebooks/wire.py @@ -0,0 +1,79 @@ +"""Decode the JSON shapes shared by the sidecar, the API and `.deepnote` files.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, cast + +from .api_types import INPUT_BLOCK_TYPES, InputBlockType +from .models import InputBlock, NotebookOutput + + +def optional_string(value: Any) -> str | None: + """Return the value when it is a string, otherwise None.""" + + return value if isinstance(value, str) else None + + +def optional_number(value: Any) -> float | int | None: + """Return the value when it is a number other than a boolean, otherwise None.""" + + return ( + value + if isinstance(value, (float, int)) and not isinstance(value, bool) + else None + ) + + +def string_tuple(value: Any) -> tuple[str, ...]: + """Return a list's items as strings, or an empty tuple for any other value.""" + + return tuple(str(item) for item in value) if isinstance(value, list) else () + + +def decode_inputs(values: Any) -> tuple[InputBlock, ...]: + """Read the sidecar's inputs, skipping any without a name or a known type.""" + + if not isinstance(values, list): + return () + return tuple( + InputBlock( + variable_name=value["variableName"], + type=cast(InputBlockType, value["type"]), + label=optional_string(value.get("label")), + value=value.get("value"), + options=string_tuple(value.get("options")), + multiple=value.get("multiple") is True, + min=optional_number(value.get("min")), + max=optional_number(value.get("max")), + step=optional_number(value.get("step")), + ) + for value in values + if isinstance(value, Mapping) + and isinstance(value.get("variableName"), str) + and value["variableName"] + and isinstance(value.get("type"), str) + and value["type"] in INPUT_BLOCK_TYPES + ) + + +def decode_block_outputs(blocks: Any, *, id_key: str) -> tuple[NotebookOutput, ...]: + """Read the outputs of a list of blocks, in block order.""" + + if not isinstance(blocks, list): + return () + outputs: list[NotebookOutput] = [] + for block in blocks: + if not isinstance(block, Mapping): + continue + block_outputs = block.get("outputs") + if not isinstance(block_outputs, list): + continue + block_id = str(block.get(id_key, "")) + block_type = optional_string(block.get("type")) + outputs.extend( + NotebookOutput(block_id=block_id, block_type=block_type, raw=output) + for output in block_outputs + if isinstance(output, Mapping) + ) + return tuple(outputs) diff --git a/deepnote_toolkit/notebooks/yaml_loader.py b/deepnote_toolkit/notebooks/yaml_loader.py new file mode 100644 index 00000000..911fd9c8 --- /dev/null +++ b/deepnote_toolkit/notebooks/yaml_loader.py @@ -0,0 +1,77 @@ +"""Load `.deepnote` YAML with the rules its writer uses.""" + +from __future__ import annotations + +import re +from typing import Any + +import yaml + +_BaseLoader: type = getattr(yaml, "CSafeLoader", yaml.SafeLoader) + + +class _DeepnoteSchemaLoader(_BaseLoader): # type: ignore[misc,valid-type] + """A safe loader for the scalar conventions used by `.deepnote` files. + + `.deepnote` files are written as YAML 1.2, where `No`, `on`, `12:30` and + `2026-08-17` are strings. PyYAML's YAML 1.1 rules read them as booleans, + numbers and dates. Leading-zero scalars intentionally remain strings, unlike + the core schema. Merge keys are treated as literal keys, not YAML 1.1 merges. + """ + + yaml_implicit_resolvers: dict[str, Any] = {} + + def construct_mapping( + self, node: yaml.MappingNode, deep: bool = False + ) -> dict[Any, Any]: + """Build a mapping, rejecting a repeated key. PyYAML would keep the last.""" + + if not isinstance(node, yaml.MappingNode): + return super().construct_mapping(node, deep=deep) + seen: set[tuple[str, str]] = set() + for key_node, _value_node in node.value: + if not isinstance(key_node, yaml.ScalarNode): + continue + key = (key_node.tag, key_node.value) + if key in seen: + raise yaml.constructor.ConstructorError( + None, + None, + f"found duplicate key {key_node.value!r}", + key_node.start_mark, + ) + seen.add(key) + return super().construct_mapping(node, deep=deep) + + +for _tag, _pattern, _first in ( + ("null", r"^(?:~|null|Null|NULL|)$", ["~", "n", "N", ""]), + ("bool", r"^(?:true|True|TRUE|false|False|FALSE)$", list("tTfF")), + # A leading zero marks a string such as a postal code. No writer emits numbers so. + ( + "int", + r"^(?:[-+]?(?:0|[1-9][0-9]*)|0o[0-7]+|0x[0-9a-fA-F]+)$", + list("-+0123456789"), + ), + ( + "float", + r"^(?:[-+]?(?:\.[0-9]+|(?:0|[1-9][0-9]*)(?:\.[0-9]*)?)(?:[eE][-+]?[0-9]+)?" + r"|[-+]?\.(?:inf|Inf|INF)|\.(?:nan|NaN|NAN))$", + list("-+0123456789."), + ), +): + _DeepnoteSchemaLoader.add_implicit_resolver( + f"tag:yaml.org,2002:{_tag}", re.compile(_pattern), _first + ) + + +def load_yaml(content: str) -> Any: + """Parse one YAML document. Raises `yaml.YAMLError` when it is malformed.""" + + loader = _DeepnoteSchemaLoader(content) + try: + return loader.get_single_data() + except (ValueError, AttributeError, TypeError) as error: + raise yaml.YAMLError(str(error)) from error + finally: + loader.dispose() diff --git a/deepnote_toolkit/streamlit/__init__.py b/deepnote_toolkit/streamlit/__init__.py new file mode 100644 index 00000000..83c8a556 --- /dev/null +++ b/deepnote_toolkit/streamlit/__init__.py @@ -0,0 +1,12 @@ +"""Helpers for Streamlit apps built on Deepnote notebooks.""" + +from .auth import CurrentUserApiTokenError, current_user_api_credentials +from .cloud_runner import StreamlitCloudRunner +from .widgets import render_inputs + +__all__ = [ + "CurrentUserApiTokenError", + "StreamlitCloudRunner", + "current_user_api_credentials", + "render_inputs", +] diff --git a/deepnote_toolkit/streamlit/auth.py b/deepnote_toolkit/streamlit/auth.py new file mode 100644 index 00000000..3f4ca458 --- /dev/null +++ b/deepnote_toolkit/streamlit/auth.py @@ -0,0 +1,239 @@ +"""Per-viewer authentication for Streamlit apps hosted by Deepnote.""" + +from __future__ import annotations + +import hashlib +import math +import os +import re +import time +from collections.abc import MutableMapping +from dataclasses import dataclass, field +from typing import Any, Protocol + +import requests +from pydantic import ValidationError +from urllib3.exceptions import LocationParseError +from urllib3.util import parse_url + +from deepnote_toolkit.get_webapp_url import ( + get_absolute_userpod_api_url, + get_project_auth_headers, +) +from deepnote_toolkit.notebooks._schemas import ViewerTokenResponse +from deepnote_toolkit.notebooks.runner import RunnerError +from deepnote_toolkit.notebooks.transport import request_json +from deepnote_toolkit.streamlit_data_apps import read_streamlit_token_from_context + +_APP_ID = r"[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}" +STREAMLIT_APP_HOST_PATTERN = re.compile(rf"^streamlit-({_APP_ID})\.", re.IGNORECASE) +STREAMLIT_APP_ID_ENV = "DEEPNOTE_STREAMLIT_APP_ID" + +_SESSION_STATE_KEY = "_deepnote_current_user_api_credentials" +_EXPIRY_MARGIN_SECONDS = 60 + + +class CurrentUserApiTokenError(RunnerError): + """Raised when a hosted app cannot obtain the current viewer's API token.""" + + +@dataclass(frozen=True) +class CurrentUserApiCredentials: + """A short-lived viewer-scoped public API credential.""" + + token: str = field(repr=False) + api_origin: str + expires_at_seconds: float + + +class StreamlitRuntime(Protocol): + """What viewer authentication asks Streamlit about the process and the thread.""" + + def has_request(self) -> bool: + """Whether this thread is running a script for a viewer.""" + + def is_worker_thread(self) -> bool: + """Whether Streamlit is running but this thread has no viewer request.""" + + def app_id(self) -> str | None: + """The hosted app's ID, from the launcher or from the request host.""" + + def viewer_cookie(self) -> str | None: + """The viewer's streamlit-token cookie.""" + + def session_state(self) -> MutableMapping[str, Any] | None: + """The viewer's session state, or None outside a script run.""" + + +class _DefaultStreamlitRuntime: + def has_request(self) -> bool: + try: + from streamlit.runtime.scriptrunner import ( # type: ignore[import-not-found] + get_script_run_ctx, + ) + except ImportError: + return False + + return get_script_run_ctx(suppress_warning=True) is not None + + def is_worker_thread(self) -> bool: + try: + from streamlit import runtime # type: ignore[import-not-found] + except ImportError: + return False + + return runtime.exists() and not self.has_request() + + def app_id(self) -> str | None: + hosted = os.environ.get(STREAMLIT_APP_ID_ENV) + if hosted is not None: + return hosted + return _read_streamlit_app_id_from_context() if self.has_request() else None + + def viewer_cookie(self) -> str | None: + return read_streamlit_token_from_context() if self.has_request() else None + + def session_state(self) -> MutableMapping[str, Any] | None: + if not self.has_request(): + return None + import streamlit as st # type: ignore[import-not-found] + + return st.session_state + + +streamlit_runtime: StreamlitRuntime = _DefaultStreamlitRuntime() + + +def current_user_api_credentials( + *, + timeout: float = 10, + session: requests.Session | None = None, + runtime: StreamlitRuntime = streamlit_runtime, +) -> CurrentUserApiCredentials: + """Exchange the viewer's cookie for short-lived public API credentials. + + The token is only valid at the returned API origin. Credentials are reused + within the current Streamlit session until shortly before they expire, and + never shared between sessions. + """ + + app_id = runtime.app_id() + if not app_id: + raise CurrentUserApiTokenError( + "Could not resolve the Deepnote Streamlit app ID." + ) + if not re.fullmatch(_APP_ID, app_id, re.IGNORECASE): + raise CurrentUserApiTokenError("The Deepnote Streamlit app ID must be a UUID.") + app_id = app_id.lower() + + viewer_token = runtime.viewer_cookie() + if not viewer_token: + raise CurrentUserApiTokenError( + "Could not read the current viewer's streamlit-token cookie." + ) + + session_state = runtime.session_state() + cache_key = (app_id, hashlib.sha256(viewer_token.encode()).hexdigest()) + if session_state is not None: + cached = session_state.get(_SESSION_STATE_KEY) + if ( + isinstance(cached, tuple) + and len(cached) == 2 + and cached[0] == cache_key + and isinstance(cached[1], CurrentUserApiCredentials) + and cached[1].expires_at_seconds - _EXPIRY_MARGIN_SECONDS > time.time() + ): + return cached[1] + + credentials = _exchange(app_id, viewer_token, timeout=timeout, session=session) + if session_state is not None: + session_state[_SESSION_STATE_KEY] = (cache_key, credentials) + return credentials + + +def _exchange( + app_id: str, + viewer_token: str, + *, + timeout: float, + session: requests.Session | None, +) -> CurrentUserApiCredentials: + http = session if session is not None else requests.Session() + try: + payload = request_json( + http, + "POST", + get_absolute_userpod_api_url(f"streamlit-apps/{app_id}/api-token"), + headers={"StreamlitToken": viewer_token, **get_project_auth_headers()}, + timeout=timeout, + ) + parsed = ViewerTokenResponse(**payload) + except RunnerError as error: + raise CurrentUserApiTokenError(str(error), transient=error.transient) from error + except ValidationError: + # Pydantic validation errors can contain the bearer token as input data. + raise CurrentUserApiTokenError( + "Viewer API-token response is missing or has invalid required fields." + ) from None + finally: + if http is not session: + http.close() + expires_at = float(parsed.expires_at_seconds) + if not math.isfinite(expires_at) or expires_at <= time.time(): + raise CurrentUserApiTokenError("Viewer API credentials have already expired.") + return CurrentUserApiCredentials( + token=parsed.token, + api_origin=_origin(parsed.api_origin), + expires_at_seconds=expires_at, + ) + + +def _read_streamlit_app_id_from_context() -> str | None: + """Resolve the app UUID from the external Streamlit request hostname. + + The exchange checks the viewer's token against this app, so a forged host + gains nothing. + """ + + try: + import streamlit as st # type: ignore[import-not-found] + except ImportError: + return None + + try: + headers = st.context.headers + except Exception: + return None + + if not headers: + return None + normalized_headers = {str(key).lower(): value for key, value in headers.items()} + for name in ("x-original-host", "host"): + host = normalized_headers.get(name) + if not isinstance(host, str): + continue + match = STREAMLIT_APP_HOST_PATTERN.match(host) + if match: + return match.group(1).lower() + return None + + +def _origin(value: str) -> str: + """Reduce `apiOrigin` to `scheme://host[:port]`, rejecting anything else in it.""" + + try: + # Use the same URL parser as Requests, including its port validation. + parts = parse_url(value) + bare = ( + parts.scheme in {"http", "https"} + and bool(parts.host) + and not re.search(r"[\s;]", parts.host) + and parts.auth is None + and not (parts.path or "").strip("/") + and not (parts.query or parts.fragment) + ) + except LocationParseError: + bare = False + if not bare: + raise CurrentUserApiTokenError("apiOrigin must be a valid HTTP(S) origin.") + return f"{parts.scheme}://{parts.netloc}" diff --git a/deepnote_toolkit/streamlit/cloud_runner.py b/deepnote_toolkit/streamlit/cloud_runner.py new file mode 100644 index 00000000..1d61c0ae --- /dev/null +++ b/deepnote_toolkit/streamlit/cloud_runner.py @@ -0,0 +1,74 @@ +"""The cloud runner for Streamlit apps hosted by Deepnote.""" + +from __future__ import annotations + +import time +from collections.abc import Callable, Mapping +from typing import Any + +import requests + +from deepnote_toolkit.notebooks.api_types import StorageMode +from deepnote_toolkit.notebooks.cloud_runner import DeepnoteCloudRunner, Sleep +from deepnote_toolkit.notebooks.credentials import DEFAULT_API_ORIGIN, TokenProvider +from deepnote_toolkit.notebooks.models import RunnerInfo +from deepnote_toolkit.notebooks.run_result import RunResult + +from .auth import StreamlitRuntime, streamlit_runtime +from .viewer_credentials import ViewerCredentials + + +class StreamlitCloudRunner: + """Run a notebook from a Streamlit app, as the viewer when Deepnote hosts it. + + The Streamlit script thread uses viewer authentication by default. Local + development requires `local=True` and an explicit token or token provider. + Deepnote hosting markers always override local credentials. Runs use read-only + project storage unless `storage_mode` says otherwise. + """ + + def __init__( + self, + notebook_id: str, + *, + local: bool = False, + token: str | None = None, + token_provider: TokenProvider | None = None, + base_url: str = DEFAULT_API_ORIGIN, + storage_mode: StorageMode | None = "readonly", + timeout: float = 600, + snapshot_timeout: float = 10, + poll_interval: float = 2, + session: requests.Session | None = None, + sleep: Sleep = time.sleep, + clock: Callable[[], float] = time.monotonic, + runtime: StreamlitRuntime = streamlit_runtime, + ): + session = session if session is not None else requests.Session() + self._runner = DeepnoteCloudRunner( + notebook_id, + credentials=ViewerCredentials( + token, + token_provider, + base_url=base_url, + timeout=min(timeout, 10), + session=session, + local=local, + runtime=runtime, + ), + storage_mode=storage_mode, + timeout=timeout, + snapshot_timeout=snapshot_timeout, + poll_interval=poll_interval, + session=session, + sleep=sleep, + clock=clock, + ) + + def info(self) -> RunnerInfo: + """Read the notebook's name and inputs using the current viewer.""" + return self._runner.info() + + def run(self, inputs: Mapping[str, Any]) -> RunResult: + """Run the notebook using the current viewer and return its outputs.""" + return self._runner.run(inputs) diff --git a/deepnote_toolkit/streamlit/viewer_credentials.py b/deepnote_toolkit/streamlit/viewer_credentials.py new file mode 100644 index 00000000..8b0d61a8 --- /dev/null +++ b/deepnote_toolkit/streamlit/viewer_credentials.py @@ -0,0 +1,78 @@ +"""API credentials of the person viewing a hosted Streamlit app.""" + +from __future__ import annotations + +import os + +import requests + +from deepnote_toolkit.notebooks.credentials import ( + DEFAULT_API_ORIGIN, + ApiCredentials, + TokenProvider, + token_credentials, +) +from deepnote_toolkit.notebooks.runner import RunnerError + +from .auth import StreamlitRuntime, current_user_api_credentials, streamlit_runtime + +_NO_REQUEST = ( + "No viewer request is available on this thread. Call the runner from the " + "Streamlit script thread." +) + + +class ViewerCredentials: + """Credentials of the current viewer when Deepnote hosts the app. + + Local Streamlit development requires `local=True` and explicit credentials. + Hosted processes and requests always use the viewer. Worker threads cannot + resolve a viewer and raise instead of using a shared token. + """ + + def __init__( + self, + token: str | None = None, + token_provider: TokenProvider | None = None, + *, + base_url: str = DEFAULT_API_ORIGIN, + timeout: float = 10, + session: requests.Session | None = None, + local: bool = False, + runtime: StreamlitRuntime = streamlit_runtime, + ): + self._local_mode = local + self._local_token_explicit = token is not None or token_provider is not None + self._local = token_credentials(token, token_provider, base_url=base_url) + self._timeout = timeout + self._session = session + self._runtime = runtime + + def __call__(self, *, timeout: float = 30) -> ApiCredentials: + """Return the viewer's credentials, or the local ones outside hosting.""" + + has_request = self._runtime.has_request() + is_hosted = ( + self._runtime.app_id() is not None + or bool(os.environ.get("DEEPNOTE_PROJECT_ID")) + or self._runtime.viewer_cookie() is not None + ) + if is_hosted or (has_request and not self._local_mode): + if not has_request: + raise RunnerError(_NO_REQUEST) + viewer = current_user_api_credentials( + timeout=min(timeout, self._timeout), + session=self._session, + runtime=self._runtime, + ) + return ApiCredentials(token=viewer.token, api_origin=viewer.api_origin) + + if self._runtime.is_worker_thread(): + raise RunnerError(_NO_REQUEST) + + if not self._local_mode or not self._local_token_explicit: + raise RunnerError( + "Viewer identity is unavailable. For local execution, set local=True and " + "pass token= or token_provider= explicitly." + ) + return self._local(timeout=timeout) diff --git a/deepnote_toolkit/streamlit/widgets.py b/deepnote_toolkit/streamlit/widgets.py new file mode 100644 index 00000000..40780e48 --- /dev/null +++ b/deepnote_toolkit/streamlit/widgets.py @@ -0,0 +1,197 @@ +"""Map Deepnote input blocks to native Streamlit widgets.""" + +from __future__ import annotations + +import calendar +import math +import re +from collections.abc import Iterable +from datetime import date, timedelta +from typing import Any + +from deepnote_toolkit.notebooks.models import InputBlock + +_RELATIVE_RANGE_MONTHS = { + "pastMonth": 1, + "past3months": 3, + "past6months": 6, + "pastYear": 12, +} + + +def render_inputs( + inputs: Iterable[InputBlock], container: Any = None, *, key_prefix: str = "deepnote" +) -> dict[str, Any]: + """Render input blocks and return API-ready values keyed by variable name. + + `container` may be `st`, `st.sidebar`, or a fake with the same widget methods + for tests. When it is omitted, Streamlit is imported lazily so parsing and API + clients work without the app extra. + """ + + if container is None: + import streamlit as st # type: ignore[import-not-found] + + container = st + + inputs = tuple(inputs) + names = [block.variable_name for block in inputs] + if any(not name for name in names): + raise ValueError("Input variable names must not be empty") + if len(names) != len(set(names)): + raise ValueError("Input variable names must be unique") + values: dict[str, Any] = {} + for input_block in inputs: + label = input_block.label or input_block.variable_name.replace("_", " ").title() + key = f"{key_prefix}:{input_block.variable_name}" + value = _render_one(container, input_block, label, key) + if value is not None: + values[input_block.variable_name] = value + return values + + +def _render_one(container: Any, input_block: InputBlock, label: str, key: str) -> Any: + """Render one input block using its saved default and constraints.""" + if input_block.type == "input-checkbox": + return container.checkbox(label, value=_as_bool(input_block.value), key=key) + + if input_block.type == "input-select": + options = list(input_block.options) + if input_block.multiple: + value = input_block.value + if not isinstance(value, list): + value = [] if value is None else [value] + defaults = [ + normalized for item in value if (normalized := str(item)) in options + ] + if len(defaults) != len(value): + container.warning( + f"{label}: saved selections are no longer available. Review the selection before running." + ) + return container.multiselect(label, options, default=defaults, key=key) + index = ( + options.index(str(input_block.value)) + if input_block.value is not None and str(input_block.value) in options + else None + ) + return container.selectbox(label, options, index=index, key=key) + + if input_block.type == "input-slider": + minimum = input_block.min if input_block.min is not None else 0 + maximum = input_block.max if input_block.max is not None else 100 + step = input_block.step if input_block.step is not None else 1 + if ( + not all(math.isfinite(n) for n in (minimum, maximum, step)) + or minimum >= maximum + or step <= 0 + ): + raise ValueError( + f"{label}: slider needs finite ordered bounds and a positive step" + ) + value = _as_number(input_block.value) + if input_block.value is not None and ( + value is None or not minimum <= value <= maximum + ): + container.warning( + f"{label}: the saved default is outside the slider's bounds. Review the value before running." + ) + value = min(max(value if value is not None else minimum, minimum), maximum) + if any(isinstance(number, float) for number in (minimum, maximum, value, step)): + minimum, maximum, value, step = ( + float(number) for number in (minimum, maximum, value, step) + ) + return container.slider( + label, min_value=minimum, max_value=maximum, value=value, step=step, key=key + ) + + if input_block.type == "input-date": + selected = _serialize_date( + container.date_input(label, value=_as_date(input_block.value), key=key) + ) + # Date blocks older than version 2 only parse a full timestamp. + is_timestamp = isinstance(input_block.value, str) and "T" in input_block.value + return f"{selected}T00:00:00.000Z" if selected and is_timestamp else selected + + if input_block.type == "input-date-range": + defaults = _as_date_range(input_block.value) + # One range picker cannot represent an open start or end independently. + if None in defaults: + return [ + _serialize_date( + container.date_input( + f"{label} ({endpoint})", value=value, key=f"{key}:{endpoint}" + ) + ) + for endpoint, value in zip(("start", "end"), defaults) + ] + selected = container.date_input(label, value=defaults, key=key) + if not isinstance(selected, (list, tuple)): + return None + serialized = [_serialize_date(value) for value in selected] + if len(serialized) == 1: + return None + return serialized if len(serialized) == 2 else ["", ""] + + if input_block.type == "input-textarea": + return container.text_area( + label, + value=str(input_block.value) if input_block.value is not None else "", + key=key, + ) + + return container.text_input( + label, + value=str(input_block.value) if input_block.value is not None else "", + key=key, + ) + + +def _as_bool(value: Any) -> bool: + """Decode checkbox defaults without treating the text false as truthy.""" + if isinstance(value, bool): + return value + return str(value).lower() in {"true", "1"} + + +def _as_number(value: Any) -> float | int | None: + """Read a saved slider value. None when it is missing or not a finite number.""" + try: + number = float(value) + except (TypeError, ValueError): + return None + if not math.isfinite(number): + return None + return int(number) if number.is_integer() else number + + +def _as_date(value: Any) -> date | None: + """Read a date or the date part of a timestamp. None leaves the widget empty.""" + + if isinstance(value, date): + return value + try: + return date.fromisoformat(str(value)[:10]) + except ValueError: + return None + + +def _as_date_range(value: Any) -> tuple[date | None, ...]: + """Resolve a range; None marks an open endpoint and () an entirely empty range.""" + + if isinstance(value, list): + dates = tuple(_as_date(item) for item in value[:2]) + return dates if len(dates) == 2 and any(dates) else () + + today = date.today() + if match := re.fullmatch(r"past(\d+)days|customDays(\d+)", str(value)): + return today - timedelta(days=int(match.group(1) or match.group(2))), today + if months := _RELATIVE_RANGE_MONTHS.get(str(value)): + year, month = divmod(today.year * 12 + today.month - 1 - months, 12) + last_day = calendar.monthrange(year, month + 1)[1] + return date(year, month + 1, min(today.day, last_day)), today + return () + + +def _serialize_date(value: Any) -> str: + """Encode a chosen date, leaving an empty widget empty.""" + return value.isoformat() if isinstance(value, date) else "" diff --git a/deepnote_toolkit/streamlit_data_apps.py b/deepnote_toolkit/streamlit_data_apps.py index 99fadeb3..f989e61f 100644 --- a/deepnote_toolkit/streamlit_data_apps.py +++ b/deepnote_toolkit/streamlit_data_apps.py @@ -61,7 +61,7 @@ def __init__( self.integration_name = integration_name -def _read_streamlit_token_from_context() -> Optional[str]: +def read_streamlit_token_from_context() -> Optional[str]: """Read the ``streamlit-token`` cookie from the active Streamlit context. Returns ``None`` if Streamlit is not installed, no script run is active, or the cookie @@ -121,7 +121,7 @@ def get_federated_auth_token( if not integration_id: raise StreamlitFederatedAuthError("integration_id is required.") - token = streamlit_token or _read_streamlit_token_from_context() + token = streamlit_token or read_streamlit_token_from_context() if not token: raise StreamlitFederatedAuthError( "Could not read the `streamlit-token` cookie from the Streamlit context. " diff --git a/docs/streamlit-apps.md b/docs/streamlit-apps.md new file mode 100644 index 00000000..279f0356 --- /dev/null +++ b/docs/streamlit-apps.md @@ -0,0 +1,101 @@ +# Build a Streamlit app from a Deepnote notebook + +Install `deepnote-toolkit` and `streamlit`, export your notebook as a `.deepnote` +file, and put it next to your app. Use the same notebook ID when reading inputs +and running the notebook, especially in projects with several notebooks. + +```python +import streamlit as st +from deepnote_toolkit.notebooks import DeepnoteDocument, RunnerError +from deepnote_toolkit.streamlit import StreamlitCloudRunner, render_inputs + +notebook_id = "your-notebook-id" +document = DeepnoteDocument.load("report.deepnote", notebook_id=notebook_id) +runner = StreamlitCloudRunner(notebook_id) +values = render_inputs(document.inputs, st.sidebar) + +if st.button("Run"): + try: + result = runner.run(values) + if not result.success: + st.error(result.error or "The run failed.") + elif result.snapshot_status == "pending": + st.info("The run finished, but its outputs are not available yet.") + elif (table := result.first_dataframe()) is not None: + st.dataframe(table.records(include_index=False)) + else: + st.write(result.text()) + except RunnerError as error: + st.error(str(error)) +``` + +## Authentication and local development + +On Deepnote, `StreamlitCloudRunner` runs the notebook with the current viewer's +permissions. If the app ID, viewer cookie or token exchange is unavailable, the +call fails instead of falling back to an owner token. Call it on the Streamlit +script thread, not a worker thread. + +For a locally hosted Streamlit app, opt into local credentials explicitly: + +```python +import os + +runner = StreamlitCloudRunner( + notebook_id, local=True, token=os.environ["DEEPNOTE_TOKEN"] +) +``` + +`token_provider=` can supply a renewable token instead. On Deepnote, `local=True` +and explicit tokens are ignored and the viewer is used. + +To call other API endpoints as the viewer, `current_user_api_credentials()` +returns the viewer's short-lived token and the API origin it is valid at. It +raises `CurrentUserApiTokenError` outside a hosted request. + +For Python code outside Streamlit, use `DeepnoteCloudRunner` from +`deepnote_toolkit.notebooks`. It accepts `token=`, `token_provider=`, or +`DEEPNOTE_TOKEN`. For a local `@deepnote/local-runner` sidecar, use +`DeepnoteLocalRunner(base_url="http://127.0.0.1:8787")`. + +## Inputs and outputs + +`render_inputs()` keeps saved defaults, including `0` and `False`. A select +without a valid saved choice starts empty. Unselected single selects and +incomplete date-range selections are left out of the returned dictionary, so +disable your Run button until the required values are present. An omitted input +runs with the notebook's saved value. A saved open-ended date range renders as +separate start and end fields. Stale multi-select choices and slider defaults +outside the bounds are adjusted with a warning. Invalid slider bounds, empty +variable names and duplicate variable names raise `ValueError`. File inputs render +as text paths; this helper does not upload files. + +`runner.info().matches_inputs(document.inputs)` compares static input definitions: +unique names, types, single/multiple selection, options, and slider bounds/steps. +It is a drift check, not a guarantee that every submitted value will be accepted. +Options populated from a variable cannot be checked against the saved file. + +Cloud results contain outputs from the executed notebook. `result.text()` returns +text and `result.first_dataframe()` returns the first table, if present. Tables +contain a preview page; check `row_count` and `is_truncated` before treating them as +complete data. Non-numeric cells, including booleans, can arrive as strings. +`DeepnoteDocument.load("report.snapshot.deepnote")` can display saved outputs +without network access. + +## Execution settings + +Streamlit runs use `storage_mode="readonly"`: the notebook can read the project's +files but not change them. Pass `storage_mode="read_write"` when the app needs to +write them. `DeepnoteCloudRunner` leaves the choice to the API. + +`timeout` (600 seconds by default) is the polling budget, including time spent +creating the run, obtaining credentials and reading outputs. Network operations +or custom credential providers can overrun it. Expiry stops polling but does not +cancel the notebook. Outputs can arrive after the run finishes; `snapshot_timeout` +(10 seconds) limits the remaining wait for them within that budget. A result whose +`snapshot_status` is still `pending` has no outputs yet. + +Pass `session=requests.Session()` to configure proxies or HTTP adapters. On +`DeepnoteCloudRunner`, `credentials=` accepts any callable that takes a `timeout` +keyword and returns `ApiCredentials(token=..., api_origin=...)`. Supported names +are listed in each package's `__all__`. diff --git a/installer/module/server_process.py b/installer/module/server_process.py index 89a43407..e9e7f63d 100644 --- a/installer/module/server_process.py +++ b/installer/module/server_process.py @@ -14,7 +14,13 @@ class ServerProcess: """A class to manage a server process.""" - def __init__(self, command: str, cwd: Optional[str] = None): + def __init__( + self, + command: str, + cwd: Optional[str] = None, + *, + env: Optional[dict[str, str]] = None, + ) -> None: """ Initialize the ServerProcess with the given command. @@ -22,6 +28,7 @@ def __init__(self, command: str, cwd: Optional[str] = None): """ self.command = command self.cwd = cwd + self.env = dict(env or {}) self.process = None self.stdout_thread = None self.stderr_thread = None @@ -53,7 +60,7 @@ def _start(self, retries: int = 3, delay: float = 0.2) -> subprocess.Popen: :return: The started process. :raises Exception: If the process fails to start after all retries. """ - env = os.environ.copy() + env = {**os.environ, **self.env} env["PYTHONUNBUFFERED"] = "1" attempt = 0 diff --git a/installer/module/streamlit.py b/installer/module/streamlit.py index c20c1aee..8debe7e7 100644 --- a/installer/module/streamlit.py +++ b/installer/module/streamlit.py @@ -3,6 +3,7 @@ import json import logging import os +import shlex import urllib.request from typing import List @@ -120,7 +121,10 @@ def start_streamlit_servers( processes.append( venv.start_server( - f"streamlit run '{entrypoint_path}' {arg_str}", cwd=directory_path + f"streamlit run {shlex.quote(entrypoint_path)} {arg_str}", + cwd=directory_path, + # The toolkit reads the app ID to run notebooks as the app's viewer. + env={"DEEPNOTE_STREAMLIT_APP_ID": str(app.get("id") or "")}, ) ) except Exception as e: diff --git a/installer/module/virtual_environment.py b/installer/module/virtual_environment.py index c2ccb6c3..edf74ae6 100644 --- a/installer/module/virtual_environment.py +++ b/installer/module/virtual_environment.py @@ -64,7 +64,13 @@ def execute(self, command: str) -> str: result = self._run_command(full_command, shell=True) return result.stdout - def start_server(self, command: str, cwd: Optional[str] = None) -> ServerProcess: + def start_server( + self, + command: str, + cwd: Optional[str] = None, + *, + env: Optional[dict[str, str]] = None, + ) -> ServerProcess: """ Start a server process using the virtual environment. @@ -73,7 +79,7 @@ def start_server(self, command: str, cwd: Optional[str] = None) -> ServerProcess :raises Exception: If the server fails to start. """ full_command = f". {self.activate_file_path} && {command}" - server_proc = ServerProcess(full_command, cwd=cwd) + server_proc = ServerProcess(full_command, cwd=cwd, env=env) # Start the server internally and handle any startup errors server_proc.start() diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index 5f9179d8..109ac8e0 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -1,11 +1,15 @@ """Pytest configuration and fixtures for unit tests.""" import os +import sys import tempfile -from typing import Generator +from typing import TYPE_CHECKING, Generator import pytest +if TYPE_CHECKING: + from streamlit.testing.v1 import AppTest + @pytest.fixture(autouse=True, scope="session") def apply_patches() -> None: @@ -65,3 +69,13 @@ def test_log_directory() -> Generator[str, None, None]: os.environ.pop("DEEPNOTE_PATHS__LOG_DIR", None) else: os.environ["DEEPNOTE_PATHS__LOG_DIR"] = original_log_dir + + +@pytest.fixture +def streamlit_app_test(monkeypatch: pytest.MonkeyPatch) -> "type[AppTest]": + """Restore the main module that Streamlit replaces while executing an app.""" + pytest.importorskip("streamlit") + from streamlit.testing.v1 import AppTest + + monkeypatch.setitem(sys.modules, "__main__", sys.modules["__main__"]) + return AppTest diff --git a/tests/unit/helpers/notebook_api.py b/tests/unit/helpers/notebook_api.py new file mode 100644 index 00000000..025ef514 --- /dev/null +++ b/tests/unit/helpers/notebook_api.py @@ -0,0 +1,56 @@ +"""Reusable HTTP and clock fixtures for the notebook clients.""" + +import json +from typing import Any + +import requests +import responses + + +class Clock: + def __init__(self): + self.now = 0.0 + self.sleeps = [] + + def __call__(self): + return self.now + + def sleep(self, seconds): + self.sleeps.append(seconds) + self.now += seconds + + +def create_run_response(status: str = "success", **fields: Any) -> dict[str, Any]: + """Build the run identity returned when creating a run.""" + return {"runId": "run-1", "status": status, **fields} + + +def run_response(status: str = "success", **fields: Any) -> dict[str, Any]: + """Build GET run details with the required snapshot lifecycle status.""" + return { + "runId": "run-1", + "status": status, + "snapshotStatus": ( + "pending" if status in {"pending", "running"} else "available" + ), + **fields, + } + + +def session(): + http = requests.Session() + http.trust_env = False + return http + + +def body(call): + return json.loads(call.request.body) + + +def add_run(http, payload, *, create=False, origin="https://api.deepnote.com"): + """Register the flat POST /v2/runs response or the nested GET /v2/runs/{id} one.""" + if create: + http.add(responses.POST, origin + "/v2/runs", json=payload) + else: + path = "/v2/runs/run-1?snapshotDelivery=blocks" + http.add(responses.GET, origin + path, json={"run": payload}) diff --git a/tests/unit/helpers/streamlit_runtime.py b/tests/unit/helpers/streamlit_runtime.py new file mode 100644 index 00000000..747ef1b8 --- /dev/null +++ b/tests/unit/helpers/streamlit_runtime.py @@ -0,0 +1,28 @@ +"""A Streamlit runtime whose answers a test sets directly.""" + +from dataclasses import dataclass, field +from typing import Any, Optional + + +@dataclass +class FakeStreamlitRuntime: + request: bool = True + worker: bool = False + app: Optional[str] = None + cookie: Optional[str] = None + state: Optional[dict[str, Any]] = field(default_factory=dict) + + def has_request(self) -> bool: + return self.request + + def is_worker_thread(self) -> bool: + return self.worker + + def app_id(self) -> Optional[str]: + return self.app + + def viewer_cookie(self) -> Optional[str]: + return self.cookie + + def session_state(self) -> Optional[dict[str, Any]]: + return self.state diff --git a/tests/unit/test_deepnote_streamlit_auth.py b/tests/unit/test_deepnote_streamlit_auth.py new file mode 100644 index 00000000..8e201747 --- /dev/null +++ b/tests/unit/test_deepnote_streamlit_auth.py @@ -0,0 +1,249 @@ +import sys +import time +from types import SimpleNamespace + +import pytest +import requests +import responses + +from deepnote_toolkit.streamlit import auth +from tests.unit.helpers.notebook_api import session +from tests.unit.helpers.streamlit_runtime import FakeStreamlitRuntime + +APP_ID = "3853c7f5-2048-4b57-946d-6c5592c3317e" +TOKEN_URL = f"http://localhost:19456/userpod-api/streamlit-apps/{APP_ID}/api-token" + + +@pytest.fixture +def runtime(): + return FakeStreamlitRuntime(app=APP_ID, cookie="cookie") + + +@pytest.fixture +def http(): + with responses.RequestsMock() as mock: + yield mock + + +def credentials(http_session, runtime): + return auth.current_user_api_credentials(session=http_session, runtime=runtime) + + +def payload(**overrides): + return { + "token": "viewer", + "apiOrigin": "https://api.deepnote-staging.com/", + "expiresAtSeconds": time.time() + 900, + **overrides, + } + + +def test_exchange_uses_cookie_and_reuses_credentials_only_in_same_session( + http, runtime +): + http.post(TOKEN_URL, json=payload()) + transport = session() + first = credentials(transport, runtime) + assert credentials(transport, runtime) is first + assert ( + first.token == "viewer" + and first.api_origin == "https://api.deepnote-staging.com" + ) + assert len(http.calls) == 1 + assert http.calls[0].request.headers["StreamlitToken"] == "cookie" + assert "Authorization" not in http.calls[0].request.headers + runtime.state.clear() + assert credentials(transport, runtime) is not first + assert len(http.calls) == 2 + + +def test_changed_cookie_or_expiry_refreshes_credentials(http, runtime): + http.post(TOKEN_URL, json=payload(expiresAtSeconds=time.time() + 30)) + http.post(TOKEN_URL, json=payload(token="second")) + http.post(TOKEN_URL, json=payload(token="third")) + transport = session() + assert credentials(transport, runtime).token == "viewer" + assert credentials(transport, runtime).token == "second" + runtime.cookie = "changed" + assert credentials(transport, runtime).token == "third" + + +@pytest.mark.parametrize( + "cache_shape", + ["empty", "short", "long", "none", "text", "mapping", "old_type", "list"], +) +def test_malformed_cached_credentials_are_refreshed( + http: responses.RequestsMock, + runtime: FakeStreamlitRuntime, + cache_shape: str, +) -> None: + """Replace malformed or stale session state with a fresh credential exchange.""" + http.post(TOKEN_URL, json=payload(token="old")) + http.post(TOKEN_URL, json=payload(token="fresh")) + with session() as transport: + initial = credentials(transport, runtime) + assert runtime.state is not None + key, _ = runtime.state[auth._SESSION_STATE_KEY] + malformed = { + "empty": (), + "short": (key,), + "long": (key, initial, "extra"), + "none": (key, None), + "text": (key, "old-token"), + "mapping": (key, {"expires_at_seconds": time.time() + 900}), + "old_type": (key, SimpleNamespace(**vars(initial))), + "list": [key, initial], + } + runtime.state[auth._SESSION_STATE_KEY] = malformed[cache_shape] + + refreshed = credentials(transport, runtime) + assert refreshed.token == "fresh" + assert runtime.state[auth._SESSION_STATE_KEY] == (key, refreshed) + assert credentials(transport, runtime) is refreshed + assert len(http.calls) == 2 + + +@pytest.mark.parametrize("value", ["bad/path", "../apps", "", "x?query", "x#fragment"]) +def test_app_id_is_validated_before_network(http, runtime, value): + runtime.app = value + with pytest.raises(auth.CurrentUserApiTokenError): + credentials(session(), runtime) + assert not http.calls + + +@pytest.mark.parametrize( + "overrides", + [ + {"token": ""}, + {"token": 1}, + {"expiresAtSeconds": True}, + {"expiresAtSeconds": 0}, + {"expiresAtSeconds": "99999999999"}, + {"apiOrigin": "https://user:pass@example.com"}, + {"apiOrigin": "https://example.com/path"}, + {"apiOrigin": "https://example.com/?next=1"}, + {"apiOrigin": "ftp://example.com"}, + {"apiOrigin": "https://[::1"}, + {"apiOrigin": "https://example.com;/"}, + {"apiOrigin": "https://example.com;"}, + {"apiOrigin": "https://exam;ple.com/"}, + {"apiOrigin": "https://example.com:8443;/"}, + {"apiOrigin": "https://example.com:invalid/"}, + {"apiOrigin": "https://example.com:65536/"}, + {"apiOrigin": "https://exa mple.com/"}, + {"apiOrigin": "https://example.com\t/"}, + {"apiOrigin": "https://example.com\\other/"}, + ], +) +def test_malformed_credentials_are_not_cached(http, runtime, overrides): + http.post(TOKEN_URL, json=payload(**overrides)) + with pytest.raises(auth.CurrentUserApiTokenError): + credentials(session(), runtime) + assert runtime.state == {} + + +@pytest.mark.parametrize("suffix", ["", "/", "/?", "/#", "?#"]) +@pytest.mark.parametrize( + "origin,expected", + [ + ("HTTPS://api.example.com:8443", "https://api.example.com:8443"), + ("http://localhost:8080", "http://localhost:8080"), + ("https://[::1]:8443", "https://[::1]:8443"), + ], +) +def test_api_origin_is_reduced_to_scheme_and_host( + http: responses.RequestsMock, + runtime: FakeStreamlitRuntime, + suffix: str, + origin: str, + expected: str, +) -> None: + """Preserve hosts and ports while removing empty URL components.""" + http.post(TOKEN_URL, json=payload(apiOrigin=origin + suffix)) + assert credentials(session(), runtime).api_origin == expected + + +@pytest.mark.parametrize( + "status,transient", [(401, False), (403, False), (429, True), (503, True)] +) +def test_exchange_preserves_server_reason_and_retry_classification( + http, runtime, status, transient +): + http.post( + TOKEN_URL, + status=status, + json={"message": "API access is not available for this app"}, + ) + with pytest.raises( + auth.CurrentUserApiTokenError, match="API access is not available" + ) as exc: + credentials(session(), runtime) + assert exc.value.transient is transient + assert len(http.calls) == 1 + + +@pytest.mark.parametrize( + "failure", [requests.Timeout(), requests.ConnectionError("closed")] +) +def test_exchange_network_failures_are_transient(http, runtime, failure): + http.post(TOKEN_URL, body=failure) + with pytest.raises(auth.CurrentUserApiTokenError) as exc: + credentials(session(), runtime) + assert exc.value.transient + + +def test_exchange_never_follows_redirects_or_exposes_html(http, runtime): + http.post(TOKEN_URL, status=302, headers={"Location": "https://other.example"}) + with pytest.raises(auth.CurrentUserApiTokenError, match="Refused a redirect"): + credentials(session(), runtime) + http.replace(responses.POST, TOKEN_URL, status=502, body="private") + with pytest.raises(auth.CurrentUserApiTokenError) as exc: + credentials(session(), runtime) + assert "private" not in str(exc.value) + + +@pytest.mark.parametrize( + "headers,expected", + [ + ({"Host": f"streamlit-{APP_ID}.example"}, APP_ID), + ( + {"Host": "localhost", "X-Original-Host": f"streamlit-{APP_ID}.example"}, + APP_ID, + ), + ({"Host": "localhost:8501"}, None), + ({}, None), + ], +) +def test_host_id_resolution(monkeypatch, headers, expected): + monkeypatch.setitem( + sys.modules, + "streamlit", + SimpleNamespace(context=SimpleNamespace(headers=headers)), + ) + assert auth._read_streamlit_app_id_from_context() == expected + + +def test_session_state_lookup_suppresses_missing_context_warning(monkeypatch): + calls = [] + + def get_ctx(*, suppress_warning): + calls.append(suppress_warning) + return None + + monkeypatch.setitem(sys.modules, "streamlit", SimpleNamespace()) + monkeypatch.setitem( + sys.modules, + "streamlit.runtime.scriptrunner", + SimpleNamespace(get_script_run_ctx=get_ctx), + ) + assert auth.streamlit_runtime.session_state() is None + assert calls == [True] + + +def test_credential_validation_traceback_does_not_expose_bearer(http, runtime): + import traceback + + http.post(TOKEN_URL, json=payload(token={"secret": "private-token"})) + with pytest.raises(auth.CurrentUserApiTokenError) as exc: + credentials(session(), runtime) + assert "private-token" not in "".join(traceback.format_exception(exc.value)) diff --git a/tests/unit/test_deepnote_streamlit_cloud_runner.py b/tests/unit/test_deepnote_streamlit_cloud_runner.py new file mode 100644 index 00000000..67b4cd15 --- /dev/null +++ b/tests/unit/test_deepnote_streamlit_cloud_runner.py @@ -0,0 +1,251 @@ +import time +from pathlib import Path +from typing import Any + +import pytest +import requests +import responses + +from deepnote_toolkit.notebooks import RunnerError +from deepnote_toolkit.streamlit import StreamlitCloudRunner, auth +from tests.unit.helpers.notebook_api import ( + Clock, + add_run, + body, + create_run_response, + run_response, + session, +) +from tests.unit.helpers.streamlit_runtime import FakeStreamlitRuntime + +APP_ID = "11111111-2222-3333-4444-555555555555" +TOKEN_URL = f"http://localhost:19456/userpod-api/streamlit-apps/{APP_ID}/api-token" + + +def viewer_token(**overrides): + return { + "token": "viewer", + "apiOrigin": "https://api.deepnote.com", + "expiresAtSeconds": time.time() + 900, + **overrides, + } + + +@pytest.fixture(autouse=True) +def env(monkeypatch): + monkeypatch.delenv("DEEPNOTE_STREAMLIT_APP_ID", raising=False) + monkeypatch.delenv("DEEPNOTE_PROJECT_ID", raising=False) + monkeypatch.setenv("DEEPNOTE_TOKEN", "owner-token") + + +@pytest.fixture +def http(): + with responses.RequestsMock() as mock: + yield mock + + +def hosted(**overrides): + return FakeStreamlitRuntime( + **{"app": APP_ID, "cookie": "viewer-cookie", **overrides} + ) + + +@pytest.mark.parametrize( + "explicit", [{}, {"token": "owner"}, {"token_provider": lambda: "owner"}] +) +def test_script_thread_without_local_mode_fails_closed(http, explicit): + runner = StreamlitCloudRunner( + "n", session=session(), runtime=FakeStreamlitRuntime(), **explicit + ) + with pytest.raises(RunnerError, match="app ID"): + runner.run({}) + assert len(http.calls) == 0 + + +@pytest.mark.parametrize("local", [False, True]) +def test_hosted_run_uses_viewer_and_readonly_even_with_explicit_owner_token( + http, local +): + http.post( + TOKEN_URL, + json=viewer_token(apiOrigin="https://api.deepnote-staging.com"), + ) + add_run( + http, + create_run_response("running"), + create=True, + origin="https://api.deepnote-staging.com", + ) + add_run( + http, + run_response(snapshotBlocks=[]), + origin="https://api.deepnote-staging.com", + ) + runner = StreamlitCloudRunner( + "n", + token="owner", + base_url="https://wrong.example", + local=local, + session=session(), + runtime=hosted(), + ) + assert runner.run({}).success + assert http.calls[0].request.headers["StreamlitToken"] == "viewer-cookie" + assert http.calls[1].request.headers["Authorization"] == "Bearer viewer" + assert body(http.calls[1])["detachedRunStorageMode"] == "readonly" + + +@pytest.mark.parametrize("explicit", [{}, {"local": True, "token": "owner"}]) +def test_project_marker_fails_closed_without_an_app_id(http, monkeypatch, explicit): + monkeypatch.setenv("DEEPNOTE_PROJECT_ID", "project") + runner = StreamlitCloudRunner( + "n", session=session(), runtime=FakeStreamlitRuntime(), **explicit + ) + with pytest.raises(RunnerError, match="app ID"): + runner.run({}) + assert len(http.calls) == 0 + + +@pytest.mark.parametrize("ambient_auth", ["netrc", "session"]) +def test_resolved_viewer_token_overrides_requests_auth( + http: responses.RequestsMock, + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + ambient_auth: str, +) -> None: + """Keep the viewer bearer authoritative over ambient Requests authentication.""" + http.post(TOKEN_URL, json=viewer_token()) + add_run(http, create_run_response("running"), create=True) + add_run(http, run_response(snapshotBlocks=[])) + transport = requests.Session() + if ambient_auth == "netrc": + netrc = tmp_path / "credentials.netrc" + netrc.write_text("machine api.deepnote.com login unrelated password dummy\n") + monkeypatch.setenv("NETRC", str(netrc)) + else: + + def owner_auth(request: requests.PreparedRequest) -> requests.PreparedRequest: + request.headers["Authorization"] = "Bearer owner" + return request + + transport.auth = owner_auth + + with transport: + result = StreamlitCloudRunner("n", session=transport, runtime=hosted()).run({}) + + assert result.success + assert http.calls[1].request.headers["Authorization"] == "Bearer viewer" + assert transport.trust_env is True + + +@pytest.mark.parametrize("marker", ["", "invalid", APP_ID]) +def test_malformed_app_marker_never_falls_back(http, marker): + if marker == APP_ID: + http.post( + TOKEN_URL, + status=403, + json={"message": "API access is not available for this app"}, + ) + runner = StreamlitCloudRunner( + "n", token="owner", local=True, session=session(), runtime=hosted(app=marker) + ) + with pytest.raises(RunnerError): + runner.info() + assert all(c.request.url == TOKEN_URL for c in http.calls) + + +@pytest.mark.parametrize( + "explicit", [{}, {"token": "owner"}, {"token_provider": lambda: "owner"}] +) +def test_worker_thread_fails_closed(http, explicit): + runtime = FakeStreamlitRuntime(request=False, worker=True) + with pytest.raises(RunnerError, match="No viewer request"): + StreamlitCloudRunner("n", session=session(), runtime=runtime, **explicit).info() + assert not http.calls + + +@pytest.mark.parametrize("script", [False, True]) +def test_local_streamlit_requires_explicit_opt_in_and_token( + http: responses.RequestsMock, script: bool +) -> None: + """Explicit local credentials work with or without an active Streamlit script.""" + runtime = FakeStreamlitRuntime(request=script) + http.get( + "https://api.deepnote.com/v2/notebooks/n", json={"notebook": {"name": "N"}} + ) + runner = StreamlitCloudRunner( + "n", local=True, token="local", session=session(), runtime=runtime + ) + assert runner.info().notebook == "N" + assert http.calls[0].request.headers["Authorization"] == "Bearer local" + with pytest.raises(RunnerError, match="explicitly"): + StreamlitCloudRunner("n", local=True, session=session(), runtime=runtime).info() + + +@pytest.mark.parametrize("explicit", [{}, {"token": "owner"}, {"local": True}]) +def test_bare_python_requires_explicit_local_credentials( + http: responses.RequestsMock, explicit: dict[str, Any] +) -> None: + """The Streamlit adapter cannot use an ambient owner token outside the runtime.""" + runtime = FakeStreamlitRuntime(request=False) + with pytest.raises(RunnerError, match=r"local=True.*explicitly"): + StreamlitCloudRunner("n", session=session(), runtime=runtime, **explicit).info() + assert not http.calls + + +def test_transient_exchange_failure_during_poll_is_retried(http): + http.post(TOKEN_URL, json=viewer_token()) + http.post(TOKEN_URL, status=503) + http.post(TOKEN_URL, json=viewer_token()) + add_run(http, create_run_response("running"), create=True) + add_run(http, run_response(snapshotBlocks=[])) + clock = Clock() + runner = StreamlitCloudRunner( + "n", + session=session(), + clock=clock, + sleep=clock.sleep, + runtime=hosted(state=None), + ) + assert runner.run({}).success + assert len(http.calls) == 5 + + +def test_real_streamlit_script_and_worker_keep_viewer_identity( + monkeypatch, http, streamlit_app_test +): + monkeypatch.setenv("DEEPNOTE_STREAMLIT_APP_ID", APP_ID) + monkeypatch.setattr(auth, "read_streamlit_token_from_context", lambda: "cookie") + http.post(TOKEN_URL, json=viewer_token()) + add_run(http, create_run_response("running"), create=True) + add_run(http, run_response(snapshotBlocks=[])) + + def app(): + import threading + + import streamlit as st + + from deepnote_toolkit.notebooks import RunnerError + from deepnote_toolkit.streamlit import StreamlitCloudRunner + + runner = StreamlitCloudRunner("n", token="owner", local=True) + st.session_state["success"] = runner.run({}).success + errors = [] + + def worker(): + try: + runner.run({}) + except RunnerError as error: + errors.append(str(error)) + + thread = threading.Thread(target=worker) + thread.start() + thread.join(timeout=5) + st.session_state["worker_errors"] = errors + + at = streamlit_app_test.from_function(app).run() + assert not at.exception + assert at.session_state["success"] is True + assert "No viewer request" in at.session_state["worker_errors"][0] + assert len(http.calls) == 3 + assert http.calls[1].request.headers["Authorization"] == "Bearer viewer" diff --git a/tests/unit/test_deepnote_streamlit_widgets.py b/tests/unit/test_deepnote_streamlit_widgets.py new file mode 100644 index 00000000..b65de96c --- /dev/null +++ b/tests/unit/test_deepnote_streamlit_widgets.py @@ -0,0 +1,461 @@ +from datetime import date +from typing import TYPE_CHECKING, Any + +import pytest + +from deepnote_toolkit.notebooks import InputBlock +from deepnote_toolkit.streamlit import render_inputs + +if TYPE_CHECKING: + from streamlit.testing.v1 import AppTest + + +class FakeContainer: + """Return widget defaults without starting Streamlit.""" + + def warning(self, message: str) -> None: + """Ignore warnings unless a test installs a recording callback.""" + pass + + def checkbox(self, _label: str, **kwargs: Any) -> Any: + """Return the configured checkbox value.""" + return kwargs["value"] + + def multiselect(self, _label: str, _options: list[str], **kwargs: Any) -> Any: + """Return the configured selection list.""" + return kwargs["default"] + + def selectbox(self, _label: str, options: list[str], **kwargs: Any) -> Any: + """Return the selected option, preserving an empty selection.""" + return options[kwargs["index"]] if kwargs["index"] is not None else None + + def slider(self, _label: str, **kwargs: Any) -> Any: + """Record or return the configured slider value.""" + return kwargs["value"] + + def date_input(self, _label: str, **kwargs: Any) -> Any: + """Return the dates supplied by the test container.""" + return kwargs["value"] + + def text_area(self, _label: str, **kwargs: Any) -> Any: + """Return the configured multiline text.""" + return kwargs["value"] + + def text_input(self, _label: str, **kwargs: Any) -> Any: + """Return the configured text.""" + return kwargs["value"] + + +def test_render_inputs_maps_all_deepnote_input_types_to_api_values() -> None: + """Convert each supported widget value to its API representation.""" + inputs = [ + InputBlock("name", "input-text", "Ada"), + InputBlock("notes", "input-textarea", "Hello"), + InputBlock("enabled", "input-checkbox", True), + InputBlock("region", "input-select", "Europe", options=("All", "Europe")), + InputBlock( + "regions", + "input-select", + ["Europe"], + options=("All", "Europe"), + multiple=True, + ), + InputBlock("limit", "input-slider", "20", min=10, max=100, step=10), + InputBlock("as_of", "input-date", date(2026, 8, 17)), + InputBlock("period", "input-date-range", [date(2026, 8, 1), date(2026, 8, 17)]), + ] + + assert render_inputs(inputs, FakeContainer()) == { + "name": "Ada", + "notes": "Hello", + "enabled": True, + "region": "Europe", + "regions": ["Europe"], + "limit": 20, + "as_of": "2026-08-17", + "period": ["2026-08-01", "2026-08-17"], + } + + +def test_incomplete_date_range_is_still_valid_for_runner_contract() -> None: + """Omit incomplete ranges until the user chooses both dates.""" + + class IncompleteDateContainer(FakeContainer): + """Simulate a user selecting only the start of a range.""" + + def date_input(self, _label: str, **_kwargs: Any) -> Any: + """Return the dates supplied by the test container.""" + return (date(2026, 8, 17),) + + values = render_inputs( + [ + InputBlock( + "period", "input-date-range", [date(2026, 8, 1), date(2026, 8, 17)] + ) + ], + IncompleteDateContainer(), + ) + + assert values == {} + + +def test_slider_preserves_fractional_default_with_integer_bounds() -> None: + """Keep fractional slider defaults and consistent numeric argument types.""" + + class SliderContainer(FakeContainer): + """Record slider arguments for numeric consistency checks.""" + + slider_kwargs: dict[str, Any] + + def slider(self, _label: str, **kwargs: Any) -> Any: + """Record or return the configured slider value.""" + self.slider_kwargs = kwargs + return kwargs["value"] + + container = SliderContainer() + values = render_inputs( + [InputBlock("threshold", "input-slider", "20.5", min=10, max=30, step=0.5)], + container, + ) + + assert values == {"threshold": 20.5} + assert container.slider_kwargs == { + "min_value": 10.0, + "max_value": 30.0, + "value": 20.5, + "step": 0.5, + "key": "deepnote:threshold", + } + + +def test_multiselect_normalizes_and_filters_stale_defaults() -> None: + """Normalize saved selections and filter unavailable options.""" + values = render_inputs( + [ + InputBlock( + "regions", + "input-select", + [1, "Europe", "Missing"], + options=("1", "Europe"), + multiple=True, + ) + ], + FakeContainer(), + ) + + assert values == {"regions": ["1", "Europe"]} + + +def test_date_reads_timestamp_default_and_keeps_its_shape() -> None: + """Preserve timestamp compatibility for older date blocks.""" + defaults = [] + + class RecordingContainer(FakeContainer): + """Record defaults and simulate a changed date.""" + + def date_input(self, _label: str, **kwargs: Any) -> Any: + """Return the dates supplied by the test container.""" + defaults.append(kwargs["value"]) + return date(2026, 8, 20) + + values = render_inputs( + [ + InputBlock("legacy", "input-date", "2026-08-17T00:00:00.000Z"), + InputBlock("current", "input-date", "2026-08-17"), + ], + RecordingContainer(), + ) + + assert defaults == [date(2026, 8, 17), date(2026, 8, 17)] + assert values == {"legacy": "2026-08-20T00:00:00.000Z", "current": "2026-08-20"} + + +def test_empty_dates_stay_empty_instead_of_becoming_today() -> None: + """Leave unspecified dates empty.""" + values = render_inputs( + [ + InputBlock("as_of", "input-date", ""), + InputBlock("period", "input-date-range", ["", ""]), + ], + FakeContainer(), + ) + + assert values == {"as_of": "", "period": ["", ""]} + + +@pytest.mark.parametrize("value", [["2026-08-01", ""], ["", "2026-08-17"]]) +def test_saved_open_ended_date_range_keeps_its_chosen_endpoint( + value: list[str], +) -> None: + """Preserve the saved endpoint instead of clearing an open-ended range.""" + assert render_inputs( + [InputBlock("period", "input-date-range", value)], FakeContainer() + ) == {"period": value} + + +@pytest.mark.parametrize("value", [["2026-08-01", ""], ["", "2026-08-17"]]) +def test_real_widgets_preserve_and_edit_open_ended_date_ranges( + streamlit_app_test: "type[AppTest]", value: list[str] +) -> None: + """Keep open endpoints visible and allow completing them in real widgets.""" + + def app(value: list[str]) -> None: + """Render a saved open-ended range inside a Streamlit script.""" + import streamlit as st + + from deepnote_toolkit.notebooks import InputBlock + from deepnote_toolkit.streamlit import render_inputs + + st.session_state["values"] = render_inputs( + [InputBlock("period", "input-date-range", value)] + ) + + at = streamlit_app_test.from_function(app, args=(value,)).run() + assert not at.exception + assert at.session_state["values"] == {"period": value} + assert len(at.date_input) == 2 + missing = value.index("") + chosen = date(2026, 8, 1 if missing == 0 else 17) + at.date_input[missing].set_value(chosen).run() + assert not at.exception + assert at.session_state["values"] == {"period": ["2026-08-01", "2026-08-17"]} + + +class FrozenDate(date): + """Keep relative date calculations deterministic.""" + + @classmethod + def today(cls) -> "FrozenDate": + """Use a month end in a leap year.""" + return cls(2024, 3, 31) + + +def test_relative_date_ranges_resolve_to_concrete_dates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Resolve relative ranges and clamp dates at month boundaries.""" + monkeypatch.setattr("deepnote_toolkit.streamlit.widgets.date", FrozenDate) + + values = render_inputs( + [ + InputBlock("week", "input-date-range", "past7days"), + InputBlock("custom", "input-date-range", "customDays3"), + InputBlock("month", "input-date-range", "pastMonth"), + InputBlock("year", "input-date-range", "pastYear"), + ], + FakeContainer(), + ) + + # Mar 31 has no counterpart a month earlier and clamps to Feb 29. + assert values == { + "week": ["2024-03-24", "2024-03-31"], + "custom": ["2024-03-28", "2024-03-31"], + "month": ["2024-02-29", "2024-03-31"], + "year": ["2023-03-31", "2024-03-31"], + } + + +def test_render_inputs_runs_on_real_streamlit_widgets( + streamlit_app_test: "type[AppTest]", +) -> None: + """Exercise every widget family using Streamlit AppTest.""" + + def app() -> None: + """Render the test inputs inside a real Streamlit script.""" + import streamlit as st + + from deepnote_toolkit.notebooks import InputBlock + from deepnote_toolkit.streamlit import render_inputs + + st.session_state["values"] = render_inputs( + [ + InputBlock("name", "input-text", "Ada"), + InputBlock("enabled", "input-checkbox", True), + InputBlock("region", "input-select", "EU", options=("US", "EU")), + InputBlock( + "regions", + "input-select", + ["EU"], + options=("US", "EU"), + multiple=True, + ), + InputBlock("limit", "input-slider", "20", min=0, max=100, step=5), + InputBlock("day", "input-date", "2026-08-17"), + InputBlock("no_day", "input-date", ""), + InputBlock("span", "input-date-range", ["2026-08-01", "2026-08-17"]), + InputBlock("no_span", "input-date-range", ["", ""]), + ] + ) + + at = streamlit_app_test.from_function(app).run() + + assert not at.exception + assert at.session_state["values"] == { + "name": "Ada", + "enabled": True, + "region": "EU", + "regions": ["EU"], + "limit": 20, + "day": "2026-08-17", + "no_day": "", + "span": ["2026-08-01", "2026-08-17"], + "no_span": ["", ""], + } + + +def test_duplicate_variable_names_are_rejected() -> None: + """Reject duplicate variables before widgets overwrite their values.""" + with pytest.raises(ValueError, match="unique"): + render_inputs( + [ + InputBlock("region", "input-text", "EU"), + InputBlock("region", "input-text", "US"), + ], + FakeContainer(), + ) + + +def test_empty_variable_name_is_rejected() -> None: + """Do not render a hand-built input that cannot be submitted to the API.""" + with pytest.raises(ValueError, match="must not be empty"): + render_inputs([InputBlock("", "input-text", "value")], FakeContainer()) + + +def test_multiselect_treats_a_scalar_default_as_one_selection() -> None: + """Normalize scalar and absent multiselect defaults.""" + values = render_inputs( + [ + InputBlock( + "regions", "input-select", "EU", options=("EU", "US"), multiple=True + ), + InputBlock( + "empty", "input-select", None, options=("EU", "US"), multiple=True + ), + ], + FakeContainer(), + ) + + assert values == {"regions": ["EU"], "empty": []} + + +@pytest.mark.parametrize("kind", ["input-text", "input-textarea", "input-file"]) +@pytest.mark.parametrize("value,expected", [(0, "0"), (False, "False"), (None, "")]) +def test_falsey_text_defaults_are_preserved( + kind: str, value: Any, expected: str +) -> None: + """Keep zero and false defaults visible in text widgets.""" + assert render_inputs([InputBlock("x", kind, value)], FakeContainer()) == { + "x": expected + } + + +@pytest.mark.parametrize("value", [None, "stale"]) +def test_unselected_single_select_does_not_submit_first_option( + value: str | None, +) -> None: + """Do not submit an option that the user has not selected.""" + assert ( + render_inputs( + [InputBlock("x", "input-select", value, options=("first", "second"))], + FakeContainer(), + ) + == {} + ) + + +def test_stale_multiselect_default_warns() -> None: + """Make unavailable saved selections visible to the user.""" + warnings = [] + container = FakeContainer() + container.warning = warnings.append + values = render_inputs( + [ + InputBlock( + "x", + "input-select", + ["old", "current"], + options=("current",), + multiple=True, + ) + ], + container, + ) + assert values == {"x": ["current"]} + assert len(warnings) == 1 + + +@pytest.mark.parametrize( + "min_value,max_value,step", + [(10, 0, 1), (0, 10, 0), (0, 10, -1), (0, float("inf"), 1)], +) +def test_invalid_slider_constraints_are_rejected( + min_value: float, max_value: float, step: float +) -> None: + with pytest.raises(ValueError, match="slider"): + render_inputs( + [ + InputBlock( + "x", "input-slider", 3, min=min_value, max=max_value, step=step + ) + ], + FakeContainer(), + ) + + +@pytest.mark.parametrize( + "value,expected", [(11, 10), (-1, 0), ("bad", 0), (float("nan"), 0), (None, 0)] +) +def test_slider_default_outside_bounds_is_clamped_with_a_warning( + value: Any, expected: int +) -> None: + warnings = [] + container = FakeContainer() + container.warning = warnings.append + values = render_inputs( + [InputBlock("x", "input-slider", value, min=0, max=10, step=1)], container + ) + assert values == {"x": expected} + assert len(warnings) == (0 if value is None else 1) + + +def test_real_widgets_keep_falsey_defaults_and_require_selection( + streamlit_app_test: "type[AppTest]", +) -> None: + """Verify user interactions preserve defaults and omit partial ranges.""" + + def app() -> None: + """Render the test inputs inside a real Streamlit script.""" + import streamlit as st + + from deepnote_toolkit.notebooks import InputBlock + from deepnote_toolkit.streamlit import render_inputs + + st.session_state["values"] = render_inputs( + [ + InputBlock("zero", "input-text", 0), + InputBlock("false", "input-textarea", False), + InputBlock("choice", "input-select", "stale", options=("A", "B")), + InputBlock("period", "input-date-range", ["", ""]), + ] + ) + + at = streamlit_app_test.from_function(app).run() + assert not at.exception + assert at.text_input[0].value == "0" and at.text_area[0].value == "False" + assert "choice" not in at.session_state["values"] + at.selectbox[0].select("B").run() + assert at.session_state["values"]["choice"] == "B" + at.date_input[0].set_value((date(2026, 8, 17),)).run() + assert not at.exception + assert "period" not in at.session_state["values"] + + +def test_missing_select_value_does_not_select_literal_none_option() -> None: + """Distinguish a missing default from an option containing the word None.""" + assert ( + render_inputs( + [InputBlock("x", "input-select", None, options=("None", "EU"))], + FakeContainer(), + ) + == {} + ) diff --git a/tests/unit/test_notebooks_document.py b/tests/unit/test_notebooks_document.py new file mode 100644 index 00000000..b4834c95 --- /dev/null +++ b/tests/unit/test_notebooks_document.py @@ -0,0 +1,345 @@ +from collections.abc import Mapping +from pathlib import Path +from typing import Any + +import pytest +import responses + +from deepnote_toolkit.notebooks import ( + DeepnoteDataframe, + DeepnoteDocument, + DeepnoteLocalRunner, + InputBlock, + RunResult, +) +from deepnote_toolkit.notebooks.models import DATAFRAME_MIME, join_text +from tests.unit.helpers.notebook_api import session + + +def run_locally(payload: Mapping[str, Any]) -> RunResult: + with responses.RequestsMock() as http: + http.post("http://127.0.0.1:8787/api/run", json=payload) + return DeepnoteLocalRunner(session=session()).run({}) + + +def local_info(payload): + with responses.RequestsMock() as http: + http.get("http://127.0.0.1:8787/api/info", json=payload) + return DeepnoteLocalRunner(session=session()).info() + + +SNAPSHOT_YAML = """ +project: + name: Sales performance + notebooks: + - blocks: + - id: region-input + type: input-select + metadata: + deepnote_variable_name: region + deepnote_input_label: Region + deepnote_variable_value: Europe + deepnote_variable_options: [All, Europe] + - id: table + type: code + outputs: + - output_type: execute_result + data: + application/vnd.deepnote.dataframe.v3+json: + columns: + - name: _deepnote_index_column + - name: Revenue + rows: + - _deepnote_index_column: Europe + Revenue: 42 + - id: agent + type: agent + outputs: + - output_type: display_data + data: + text/markdown: "**Done**" +""" + + +def test_loads_inputs_and_structured_outputs(tmp_path: Path) -> None: + path = tmp_path / "sales.snapshot.deepnote" + path.write_text(SNAPSHOT_YAML, encoding="utf-8") + + snapshot = DeepnoteDocument.load(path) + + assert snapshot.project_name == "Sales performance" + assert snapshot.inputs == ( + InputBlock( + "region", + "input-select", + "Europe", + label="Region", + options=("All", "Europe"), + ), + ) + dataframe = snapshot.first_dataframe() + assert dataframe is not None + assert dataframe.records(include_index=False) == [{"Revenue": 42}] + assert snapshot.agent_text() == "**Done**" + + +def test_reads_input_metadata_from_file_and_api_shapes() -> None: + document = DeepnoteDocument( + { + "project": { + "notebooks": [ + { + "blocks": [ + { + "type": "input-slider", + "metadata": { + "deepnote_variable_name": "limit", + "deepnote_input_label": "Row limit", + "deepnote_variable_value": "20", + "deepnote_slider_min_value": 10, + "deepnote_slider_max_value": 100, + "deepnote_slider_step": 10, + }, + } + ] + } + ] + } + } + ) + info = local_info( + { + "inputs": [ + { + "variableName": "countries", + "type": "input-select", + "label": "Countries", + "value": ["Panama"], + "options": ["Panama", "Colombia"], + "multiple": True, + } + ] + } + ) + (file_input,) = document.inputs + (api_input,) = info.inputs + + assert file_input == InputBlock( + variable_name="limit", + type="input-slider", + label="Row limit", + value="20", + min=10, + max=100, + step=10, + ) + assert api_input.options == ("Panama", "Colombia") + assert api_input.multiple is True + + +def test_run_result_prefers_snapshot_outputs_and_preserves_cloud_fields() -> None: + result = run_locally( + { + "target": "cloud", + "success": True, + "runId": "run-1", + "status": "success", + "viewUrl": "https://deepnote.com/project/example", + "snapshotYaml": SNAPSHOT_YAML, + "outputs": [], + } + ) + + assert result.success is True + assert result.target == "cloud" + assert result.run_id == "run-1" + assert result.agent_text() == "**Done**" + + +def test_run_result_falls_back_to_inline_outputs_without_snapshot() -> None: + result = run_locally( + { + "target": "local", + "success": True, + "outputs": [ + { + "blockId": "code-1", + "outputs": [ + { + "output_type": "execute_result", + "data": { + DATAFRAME_MIME: { + "columns": [{"name": "value"}], + "rows": [{"value": 42}], + } + }, + } + ], + } + ], + } + ) + + dataframe = result.first_dataframe() + assert dataframe is not None + assert dataframe.records() == [{"value": 42}] + + +def test_run_result_falls_back_to_inline_outputs_for_malformed_snapshot( + caplog: pytest.LogCaptureFixture, +) -> None: + """Keep inline outputs and report the snapshot fallback without its contents.""" + result = run_locally( + { + "target": "cloud", + "runId": "run-fallback", + "success": True, + "snapshotYaml": "not: a deepnote snapshot", + "outputs": [ + { + "blockId": "code-1", + "outputs": [ + { + "output_type": "stream", + "text": "fallback output", + } + ], + } + ], + } + ) + + assert result.snapshot is None + assert result.text() == "fallback output" + assert "Could not parse the run snapshot; using inline outputs" in caplog.text + assert "run-fallback" in caplog.text and "cloud" in caplog.text + assert "not: a deepnote snapshot" not in caplog.text + + +@pytest.mark.parametrize( + ("value", "expected"), + [(["hello", " ", "world"], "hello world"), ("hello", "hello"), (None, "")], +) +def test_join_text(value: object, expected: str) -> None: + assert join_text(value) == expected + + +@pytest.mark.parametrize("content", ["hello: world", "[]", ""]) +def test_rejects_non_deepnote_yaml(content: str) -> None: + with pytest.raises(ValueError): + DeepnoteDocument.parse(content) + + +MULTI_NOTEBOOK_YAML = """ +project: + name: Sales + notebooks: + - id: notebook-a + blocks: + - type: input-text + metadata: {deepnote_variable_name: region, deepnote_variable_value: EU} + - id: notebook-b + blocks: + - type: input-text + metadata: {deepnote_variable_name: region, deepnote_variable_value: US} +""" + + +def test_notebook_id_scopes_inputs_to_one_notebook() -> None: + everything = DeepnoteDocument.parse(MULTI_NOTEBOOK_YAML) + scoped = DeepnoteDocument.parse(MULTI_NOTEBOOK_YAML, notebook_id="notebook-b") + + assert [input_block.value for input_block in everything.inputs] == ["EU", "US"] + assert scoped.inputs == (InputBlock("region", "input-text", "US"),) + + +def test_unknown_notebook_id_is_rejected() -> None: + with pytest.raises(ValueError, match="notebook-c is not in this document"): + DeepnoteDocument.parse(MULTI_NOTEBOOK_YAML, notebook_id="notebook-c") + + +def test_skips_input_blocks_of_an_unknown_type() -> None: + document = DeepnoteDocument.parse(""" +project: + notebooks: + - blocks: + - type: input-unknown + metadata: {deepnote_variable_name: mystery} + - type: input-text + metadata: {deepnote_variable_name: region} +""") + + assert document.inputs == (InputBlock("region", "input-text", None),) + + +WRITER_STYLE_YAML = """ +project: + name: Survey + notebooks: + - id: notebook-a + blocks: + - type: input-select + metadata: + deepnote_variable_name: answer + deepnote_variable_value: No + deepnote_variable_options: + - Yes + - No + - type: input-date + metadata: + deepnote_variable_name: as_of + deepnote_variable_value: 2026-08-17T00:00:00.000Z + - type: input-text + metadata: + deepnote_variable_name: time + deepnote_variable_value: 12:30 +""" + + +def test_plain_scalars_the_deepnote_writer_leaves_unquoted_stay_strings() -> None: + document = DeepnoteDocument.parse(WRITER_STYLE_YAML) + + assert document.inputs == ( + InputBlock("answer", "input-select", "No", options=("Yes", "No")), + InputBlock("as_of", "input-date", "2026-08-17T00:00:00.000Z"), + InputBlock("time", "input-text", "12:30"), + ) + + +def test_dataframe_reports_rows_beyond_the_first_page() -> None: + dataframe = DeepnoteDataframe.from_value( + {"columns": [{"name": "a"}], "rows": [{"a": 1}], "row_count": 250} + ) + whole = DeepnoteDataframe.from_value( + {"columns": [{"name": "a"}], "rows": [{"a": 1}]} + ) + + assert dataframe is not None and whole is not None + assert (dataframe.row_count, dataframe.is_truncated) == (250, True) + assert (whole.row_count, whole.is_truncated) == (1, False) + + +def test_images_decode_wrapped_base64_and_skip_invalid_data() -> None: + result = run_locally( + { + "outputs": [ + { + "blockId": "code-1", + "outputs": [ + { + "output_type": "display_data", + "data": {"image/png": "aGVs\nbG8="}, + }, + { + "output_type": "display_data", + "data": {"image/png": "not base64!"}, + }, + {"output_type": "display_data", "data": {"image/jpeg": "aGk="}}, + ], + } + ] + } + ) + + assert result.images() == [b"hello"] + assert result.images("image/jpeg") == [b"hi"] diff --git a/tests/unit/test_notebooks_runners.py b/tests/unit/test_notebooks_runners.py new file mode 100644 index 00000000..76bcdf75 --- /dev/null +++ b/tests/unit/test_notebooks_runners.py @@ -0,0 +1,768 @@ +import json +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any + +import pytest +import requests +import responses + +from deepnote_toolkit.notebooks import ( + ApiCredentials, + DeepnoteCloudRunner, + DeepnoteLocalRunner, + InputBlock, + RunnerError, + RunnerInfo, +) +from tests.unit.helpers.notebook_api import ( + Clock, + add_run, + body, + create_run_response, + run_response, + session, +) + + +@pytest.fixture +def http(): + with responses.RequestsMock() as mock: + yield mock + + +@pytest.fixture +def clock(): + return Clock() + + +@pytest.fixture +def runner(clock): + return DeepnoteCloudRunner( + "notebook-1", + token="token", + session=session(), + clock=clock, + sleep=clock.sleep, + poll_interval=0.5, + ) + + +def test_cloud_info_preserves_input_contract_and_quotes_notebook_id(http): + http.get( + "https://api.deepnote.com/v2/notebooks/a%2Fb%3Fadmin%3Dtrue", + json={ + "notebook": { + "name": "Revenue", + "inputs": [ + { + "name": "region", + "type": "input-select", + "value": "EU", + "options": ["EU", "US"], + "multiple": True, + }, + { + "name": "limit", + "type": "input-slider", + "value": "5", + "min": 1, + "max": 9, + "step": 2, + }, + {"name": "future", "type": "input-unknown"}, + ], + } + }, + ) + info = DeepnoteCloudRunner( + "a/b?admin=true", token="token", session=session() + ).info() + assert info.notebook == "Revenue" + assert info.inputs == ( + InputBlock("region", "input-select", "EU", options=("EU", "US"), multiple=True), + InputBlock("limit", "input-slider", "5", min=1, max=9, step=2), + ) + assert http.calls[0].request.headers["Authorization"] == "Bearer token" + + +def test_run_refreshes_credentials_normalizes_inputs_and_reads_only_run_blocks( + http, clock +): + add_run(http, create_run_response("pending"), create=True) + add_run(http, run_response("running")) + add_run(http, run_response(snapshotStatus="pending")) + add_run( + http, + run_response( + snapshotStatus="available", + snapshotBlocks=[ + { + "id": "b", + "type": "code", + "outputs": [{"output_type": "stream", "text": "done"}], + } + ], + snapshotContent="must not be used", + ), + ) + tokens = iter(["one", "two", "three", "four"]) + result = DeepnoteCloudRunner( + "notebook-1", + token_provider=lambda: next(tokens), + session=session(), + clock=clock, + sleep=clock.sleep, + poll_interval=0.5, + storage_mode="readonly", + ).run({"limit": 20, "enabled": False, "regions": ("EU",)}) + assert result.success and result.text() == "done" and result.snapshot is None + assert body(http.calls[0]) == { + "notebookId": "notebook-1", + "detached": True, + "inputs": {"limit": "20", "enabled": False, "regions": ["EU"]}, + "detachedRunStorageMode": "readonly", + } + assert [c.request.headers["Authorization"] for c in http.calls] == [ + f"Bearer {token}" for token in ("one", "two", "three", "four") + ] + assert clock.sleeps == [0.5, 0.5, 0.5] + + +@pytest.mark.parametrize( + "payload", + [ + {"runId": "r"}, + {"runId": "r", "status": None}, + {"runId": "r", "status": ""}, + {"runId": "r", "status": "future"}, + {"status": "success"}, + {"run": {"runId": "r", "status": "success"}}, + ], +) +def test_malformed_run_fails_without_polling(http, runner, payload): + add_run(http, payload, create=True) + with pytest.raises(RunnerError, match="invalid run response"): + runner.run({}) + assert len(http.calls) == 1 + + +@pytest.mark.parametrize( + "payload", + [ + {"runId": "run-1", "status": "future", "snapshotStatus": "pending"}, + {"runId": "run-1", "status": "success", "snapshotStatus": "future"}, + {}, + ], +) +def test_malformed_poll_stops_the_run( + http: responses.RequestsMock, + runner: DeepnoteCloudRunner, + payload: dict[str, Any], +) -> None: + """Stop polling when the API returns a response outside the run contract.""" + add_run(http, create_run_response("running"), create=True) + add_run(http, payload) + with pytest.raises(RunnerError, match="invalid run response"): + runner.run({}) + assert len(http.calls) == 2 + + +@pytest.mark.parametrize("status", ["running", "success"]) +@pytest.mark.parametrize("snapshot", [{}, {"snapshotStatus": None}]) +def test_get_run_requires_snapshot_status( + http: responses.RequestsMock, + runner: DeepnoteCloudRunner, + status: str, + snapshot: dict[str, Any], +) -> None: + """Reject missing or null GET snapshot metadata instead of returning no outputs.""" + add_run(http, create_run_response("running"), create=True) + add_run(http, {"runId": "run-1", "status": status, **snapshot}) + with pytest.raises(RunnerError, match="invalid run response"): + runner.run({}) + assert len(http.calls) == 2 + + +def test_finished_create_fetches_snapshot_metadata( + http: responses.RequestsMock, runner: DeepnoteCloudRunner +) -> None: + """A completed POST still needs GET metadata to discover its outputs.""" + add_run(http, create_run_response(), create=True) + add_run( + http, + run_response( + snapshotStatus="available", + snapshotBlocks=[ + { + "id": "b", + "type": "code", + "outputs": [{"output_type": "stream", "text": "done"}], + } + ], + ), + ) + result = runner.run({}) + assert result.success + assert result.snapshot_status == "available" + assert result.text() == "done" + assert len(http.calls) == 2 + + +@pytest.mark.parametrize( + "pending_blocks", + [ + [], + [{"id": "b", "outputs": []}], + [{"id": "b", "outputs": [{"output_type": "stream", "text": "partial"}]}], + ], +) +def test_pending_snapshot_blocks_do_not_stop_polling( + http: responses.RequestsMock, + runner: DeepnoteCloudRunner, + pending_blocks: list[dict[str, Any]], +) -> None: + """Use snapshot lifecycle status even when a pending response has blocks.""" + add_run(http, create_run_response("running"), create=True) + add_run( + http, + run_response(snapshotStatus="pending", snapshotBlocks=pending_blocks), + ) + add_run( + http, + run_response( + snapshotStatus="available", + snapshotBlocks=[ + { + "id": "b", + "outputs": [{"output_type": "stream", "text": "complete"}], + } + ], + ), + ) + result = runner.run({}) + assert result.snapshot_status == "available" + assert result.text() == "complete" + assert len(http.calls) == 3 + + +@pytest.mark.parametrize("target", ["cloud", "local"]) +def test_runner_info_skips_unnamed_inputs( + http: responses.RequestsMock, target: str +) -> None: + """Keep decoded inputs consistent with .deepnote files and valid API keys.""" + name_key = "name" if target == "cloud" else "variableName" + metadata = { + "inputs": [ + {name_key: "", "type": "input-text", "value": "unnamed"}, + {name_key: "region", "type": "input-text", "value": "EU"}, + ] + } + if target == "cloud": + http.get("https://api.deepnote.com/v2/notebooks/n", json={"notebook": metadata}) + info = DeepnoteCloudRunner("n", token="t", session=session()).info() + else: + http.get("http://127.0.0.1:8787/api/info", json=metadata) + info = DeepnoteLocalRunner(session=session()).info() + assert info.inputs == (InputBlock("region", "input-text", "EU"),) + + +def test_pending_empty_blocks_still_respect_snapshot_deadline( + http: responses.RequestsMock, clock: Clock +) -> None: + """Waiting for a pending snapshot remains bounded even when it has blocks.""" + add_run(http, create_run_response("running"), create=True) + add_run(http, run_response(snapshotStatus="pending", snapshotBlocks=[])) + runner = DeepnoteCloudRunner( + "n", + token="t", + session=session(), + snapshot_timeout=1, + poll_interval=0.25, + clock=clock, + sleep=clock.sleep, + ) + result = runner.run({}) + assert result.snapshot_status == "pending" + assert result.outputs == () + assert clock.now == 1.25 + assert len(http.calls) == 5 + + +@pytest.mark.parametrize("snapshot_status", ["unavailable", "available"]) +def test_only_pending_snapshots_are_polled(http, runner, snapshot_status): + add_run(http, create_run_response("running"), create=True) + add_run( + http, + run_response("error", snapshotStatus=snapshot_status, error="bad input"), + ) + result = runner.run({}) + assert not result.success and result.error == "bad input" + assert len(http.calls) == 2 + + +@pytest.mark.parametrize( + "failure", [429, 503, requests.ConnectionError("reset"), requests.Timeout()] +) +def test_transient_get_failure_is_retried(http, runner, failure): + add_run(http, create_run_response("running"), create=True) + kwargs = {"status": failure} if isinstance(failure, int) else {"body": failure} + http.get("https://api.deepnote.com/v2/runs/run-1?snapshotDelivery=blocks", **kwargs) + add_run(http, run_response(snapshotBlocks=[])) + assert runner.run({}).success + assert len(http.calls) == 3 + + +@pytest.mark.parametrize("status", [400, 401, 403, 404]) +def test_non_transient_poll_failure_is_not_retried(http, runner, status): + add_run(http, create_run_response("running"), create=True) + http.get( + "https://api.deepnote.com/v2/runs/run-1?snapshotDelivery=blocks", + status=status, + json={"message": "reason"}, + ) + with pytest.raises(RunnerError, match=f"HTTP {status}: reason"): + runner.run({}) + assert len(http.calls) == 2 + + +def test_post_is_never_replayed(http, runner): + http.post("https://api.deepnote.com/v2/runs", body=requests.Timeout()) + with pytest.raises(RunnerError): + runner.run({}) + assert len(http.calls) == 1 + + +def test_poll_retries_have_a_limit(http, runner): + add_run(http, create_run_response("running"), create=True) + http.get( + "https://api.deepnote.com/v2/runs/run-1?snapshotDelivery=blocks", status=503 + ) + with pytest.raises(RunnerError, match="503"): + runner.run({}) + assert len(http.calls) == 7 + + +def test_request_and_sleep_time_count_against_run_deadline(http, clock): + observed = [] + + def create(request): + observed.append(request.req_kwargs["timeout"].total) + clock.now += 3 + return 200, {}, json.dumps(create_run_response("running")) + + def poll(request): + observed.append(request.req_kwargs["timeout"].total) + clock.now += 4 + return 200, {}, json.dumps({"run": run_response("running")}) + + http.add_callback( + responses.POST, "https://api.deepnote.com/v2/runs", callback=create + ) + http.add_callback( + responses.GET, + "https://api.deepnote.com/v2/runs/run-1?snapshotDelivery=blocks", + callback=poll, + ) + runner = DeepnoteCloudRunner( + "n", + token="t", + session=session(), + timeout=10, + poll_interval=2, + clock=clock, + sleep=clock.sleep, + ) + with pytest.raises(RunnerError, match="10 seconds"): + runner.run({}) + assert observed == [10, 5] + assert clock.now == 10 + assert clock.sleeps == [2, 1] + + +@pytest.mark.parametrize("available", [False, True]) +def test_snapshot_deadline_counts_slow_requests_and_caps_timeout( + http: responses.RequestsMock, clock: Clock, available: bool +) -> None: + """Keep received outputs at the deadline without starting another request.""" + add_run(http, create_run_response(), create=True) + timeouts = [] + + def poll(request: Any) -> tuple[int, dict[str, str], str]: + """Return a snapshot as the monotonic request budget expires.""" + timeouts.append(request.req_kwargs["timeout"].total) + clock.now += 4 + payload = run_response(snapshotStatus="available" if available else "pending") + if available: + payload["snapshotBlocks"] = [ + { + "id": "b", + "type": "code", + "outputs": [{"output_type": "stream", "text": "done"}], + } + ] + return 200, {}, json.dumps({"run": payload}) + + http.add_callback( + responses.GET, + "https://api.deepnote.com/v2/runs/run-1?snapshotDelivery=blocks", + callback=poll, + ) + runner = DeepnoteCloudRunner( + "n", + token="t", + session=session(), + snapshot_timeout=5, + poll_interval=1, + clock=clock, + sleep=clock.sleep, + ) + result = runner.run({}) + assert result.snapshot_status == ("available" if available else "pending") + assert result.text() == ("done" if available else "") + assert len(http.calls) == 2 + assert timeouts == [4] + assert clock.now == 5 + + +def test_no_request_starts_after_snapshot_deadline(http, clock): + add_run(http, create_run_response(), create=True) + runner = DeepnoteCloudRunner( + "n", + token="t", + session=session(), + snapshot_timeout=0.1, + poll_interval=2, + clock=clock, + sleep=clock.sleep, + ) + assert runner.run({}).outputs == () + assert clock.sleeps == [0.1] + assert len(http.calls) == 1 + + +@pytest.mark.parametrize( + "kwargs", + [ + {"timeout": 0}, + {"timeout": float("inf")}, + {"snapshot_timeout": -1}, + {"poll_interval": 0}, + ], +) +def test_invalid_timeouts_rejected(kwargs): + with pytest.raises(ValueError): + DeepnoteCloudRunner("n", **kwargs) + + +def test_run_id_is_quoted(http, runner): + add_run(http, create_run_response("running", runId="a/b?x"), create=True) + http.get( + "https://api.deepnote.com/v2/runs/a%2Fb%3Fx?snapshotDelivery=blocks", + json={"run": run_response()}, + ) + assert runner.run({}).success + + +def test_custom_credentials_supply_origin_and_receive_budget(http): + budgets = [] + + def credentials(*, timeout): + budgets.append(timeout) + return ApiCredentials("token", "https://api.example") + + http.get("https://api.example/v2/notebooks/n", json={"notebook": {"name": "N"}}) + assert ( + DeepnoteCloudRunner("n", credentials=credentials, session=session()) + .info() + .notebook + == "N" + ) + assert budgets == [30] + + +def test_credentials_and_token_are_mutually_exclusive(): + with pytest.raises(ValueError): + DeepnoteCloudRunner("n", token="t", credentials=lambda **_: ApiCredentials("t")) + + +def test_local_runner_contract(http): + http.get( + "http://127.0.0.1:8787/api/info", + json={"notebook": "N", "runTarget": "local", "inputs": []}, + ) + http.post( + "http://127.0.0.1:8787/api/run", + json={"target": "local", "success": True, "outputs": []}, + ) + runner = DeepnoteLocalRunner(session=session()) + assert runner.info().notebook == "N" + assert runner.run({"n": 3}).success + assert body(http.calls[1]) == {"inputs": {"n": 3}} + + +@pytest.mark.parametrize("value", [None, {}, set(), b"x"]) +def test_invalid_input_is_rejected_before_sending(http, runner, value): + with pytest.raises(ValueError): + runner.run({"n": value}) + assert len(http.calls) == 0 + + +def test_error_page_does_not_leak_into_exception(http): + http.get( + "http://127.0.0.1:8787/api/info", + body="proxy internals", + status=502, + ) + with pytest.raises(RunnerError) as exc: + DeepnoteLocalRunner(session=session()).info() + assert "proxy internals" not in str(exc.value) + + +def test_real_http_redirect_never_receives_credentials(): + received = [] + + class Handler(BaseHTTPRequestHandler): + def do_GET(self): + received.append( + (self.server.server_port, self.headers.get("Authorization")) + ) + self.send_response(302) + self.send_header("Location", f"http://127.0.0.1:{other.server_port}/") + self.end_headers() + + def log_message(self, *_args): + pass + + api = HTTPServer(("127.0.0.1", 0), Handler) + other = HTTPServer(("127.0.0.1", 0), Handler) + for server in (api, other): + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + with pytest.raises(RunnerError, match="Refused a redirect"): + DeepnoteCloudRunner( + "n", + token="t", + session=session(), + base_url=f"http://127.0.0.1:{api.server_port}", + ).info() + finally: + for server in (api, other): + server.shutdown() + server.server_close() + assert received == [(api.server_port, "Bearer t")] + + +def test_runner_info_requires_matching_input_names_and_types() -> None: + info = RunnerInfo( + notebook="Revenue", + inputs=(InputBlock("region", "input-select", "All"),), + run_target="cloud", + ) + + assert info.matches_inputs([InputBlock("region", "input-select", "Europe")]) + assert not info.matches_inputs([InputBlock("market", "input-select", "Europe")]) + assert not info.matches_inputs([InputBlock("region", "input-text", "Europe")]) + + +def test_runner_info_rejects_repeated_input_names() -> None: + info = RunnerInfo( + notebook="Revenue", + inputs=(InputBlock("region", "input-text", "EU"),), + run_target="cloud", + ) + + assert not info.matches_inputs( + [ + InputBlock("region", "input-text", "EU"), + InputBlock("region", "input-text", "US"), + ] + ) + + +def test_runner_info_rejects_matching_empty_input_names() -> None: + """Matching unnamed inputs still cannot form a valid execution contract.""" + inputs = (InputBlock("", "input-text", "EU"),) + info = RunnerInfo(notebook="N", inputs=inputs, run_target="cloud") + assert not info.matches_inputs(inputs) + + +@pytest.mark.parametrize( + "changed", + [ + InputBlock("region", "input-select", "EU", options=("EU", "US"), multiple=True), + InputBlock("region", "input-select", "EU", options=("EU", "APAC")), + ], + ids=["multiple", "options"], +) +def test_runner_info_rejects_a_select_that_takes_other_values( + changed: InputBlock, +) -> None: + info = RunnerInfo( + notebook="Revenue", + inputs=(InputBlock("region", "input-select", "EU", options=("US", "EU")),), + run_target="cloud", + ) + + assert info.matches_inputs( + [InputBlock("region", "input-select", "US", options=("EU", "US"))] + ) + assert not info.matches_inputs([changed]) + + +def test_runner_info_ignores_select_options_filled_from_a_variable() -> None: + info = RunnerInfo( + notebook="Revenue", + inputs=(InputBlock("region", "input-select", "EU", options=("EU", "US")),), + run_target="cloud", + ) + + assert info.matches_inputs( + [ + InputBlock( + "region", + "input-select", + "EU", + options=("EU",), + options_from_variable=True, + ) + ] + ) + + +def test_runner_info_compares_slider_bounds_with_the_defaults_filled_in() -> None: + info = RunnerInfo( + notebook="Revenue", + inputs=(InputBlock("limit", "input-slider", "20", min=0, max=100, step=1),), + run_target="cloud", + ) + + assert info.matches_inputs([InputBlock("limit", "input-slider", "20")]) + assert not info.matches_inputs([InputBlock("limit", "input-slider", "20", max=50)]) + + +@pytest.mark.parametrize( + "changed", + [ + InputBlock("n", "input-slider", "5", min=1, max=10, step=1), + InputBlock("n", "input-slider", "5", min=0, max=9, step=1), + InputBlock("n", "input-slider", "5", min=0, max=10, step=2), + ], +) +def test_input_match_detects_slider_constraints(changed): + info = RunnerInfo( + "N", (InputBlock("n", "input-slider", "5", min=0, max=10, step=1),), "cloud" + ) + assert not info.matches_inputs([changed]) + + +def test_credential_exchange_time_reduces_http_budget(http, clock): + def credentials(*, timeout): + assert timeout == 4 + clock.now += 3 + return ApiCredentials("token") + + http.post("https://api.deepnote.com/v2/runs", json=create_run_response()) + runner = DeepnoteCloudRunner( + "n", + credentials=credentials, + timeout=4, + clock=clock, + sleep=clock.sleep, + session=session(), + ) + assert runner.run({}).success + assert http.calls[0].request.req_kwargs["timeout"].total == 1 + + +def test_expired_credential_budget_does_not_send_request(http, clock): + def credentials(*, timeout): + clock.now += timeout + return ApiCredentials("token") + + with pytest.raises(RunnerError, match="exhausted"): + DeepnoteCloudRunner( + "n", + credentials=credentials, + timeout=4, + clock=clock, + sleep=clock.sleep, + session=session(), + ).run({}) + assert not http.calls + + +@pytest.mark.parametrize("status", [403, 503]) +def test_snapshot_poll_error_policy(http, runner, status): + add_run(http, create_run_response(), create=True) + http.get( + "https://api.deepnote.com/v2/runs/run-1?snapshotDelivery=blocks", status=status + ) + if status == 403: + with pytest.raises(RunnerError, match="403"): + runner.run({}) + assert len(http.calls) == 2 + else: + add_run(http, run_response(snapshotStatus="available", snapshotBlocks=[])) + assert runner.run({}).snapshot_status == "available" + assert len(http.calls) == 3 + + +@pytest.mark.parametrize("body", ["not-json", "[]", "null"]) +def test_invalid_json_response_is_a_runner_error(http, body): + http.get("http://127.0.0.1:8787/api/info", body=body) + with pytest.raises(RunnerError, match="invalid JSON|non-object"): + DeepnoteLocalRunner(session=session()).info() + + +def test_same_origin_redirect_is_also_refused(http, runner): + http.post( + "https://api.deepnote.com/v2/runs", + status=307, + headers={"Location": "https://api.deepnote.com/other"}, + ) + with pytest.raises(RunnerError, match="Refused a redirect"): + runner.run({}) + assert len(http.calls) == 1 + + +@pytest.mark.parametrize( + "header", [None, "Authorization", "authorization", "aUtHoRiZaTiOn"] +) +def test_explicit_auth_headers_override_session_auth_regardless_of_case( + http: responses.RequestsMock, header: str | None +) -> None: + """Honor HTTP header casing while retaining session auth for unauthenticated calls.""" + from deepnote_toolkit.notebooks.transport import request_json + + def ambient_auth(request: requests.PreparedRequest) -> requests.PreparedRequest: + """Represent a session configured with a different API identity.""" + request.headers["Authorization"] = "Bearer ambient" + return request + + http.get("https://api.example/info", json={}) + with session() as transport: + transport.auth = ambient_auth + request_json( + transport, + "GET", + "https://api.example/info", + headers={header: "Bearer selected"} if header else {}, + timeout=1, + ) + + assert http.calls[0].request.headers["Authorization"] == ( + "Bearer selected" if header else "Bearer ambient" + ) + + +def test_exhausted_poll_budget_does_not_even_fetch_credentials(http): + from deepnote_toolkit.notebooks.api_client import DeepnoteApiClient + + def credentials(*, timeout): + pytest.fail("Expired requests must not fetch credentials") + + client = DeepnoteApiClient(credentials, session=session()) + with pytest.raises(RunnerError, match="deadline expired"): + client.get_run("r", timeout=0) + assert not http.calls diff --git a/tests/unit/test_notebooks_yaml_loader.py b/tests/unit/test_notebooks_yaml_loader.py new file mode 100644 index 00000000..55bf013e --- /dev/null +++ b/tests/unit/test_notebooks_yaml_loader.py @@ -0,0 +1,95 @@ +import importlib +from collections.abc import Iterator +from typing import Any + +import pytest +import yaml + +from deepnote_toolkit.notebooks import yaml_loader + + +@pytest.fixture(params=["libyaml", "pure_python"]) +def load_yaml( + request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch +) -> Iterator[Any]: + if request.param == "libyaml" and not hasattr(yaml, "CSafeLoader"): + pytest.skip("PyYAML was built without libyaml") + if request.param == "pure_python": + monkeypatch.delattr(yaml, "CSafeLoader", raising=False) + yield importlib.reload(yaml_loader).load_yaml + monkeypatch.undo() + importlib.reload(yaml_loader) + + +@pytest.mark.parametrize( + ("scalar", "expected"), + [ + ("No", "No"), + ("yes", "yes"), + ("on", "on"), + ("true", True), + ("FALSE", False), + ("12:30", "12:30"), + ("1_000", "1_000"), + ("2026-08-17", "2026-08-17"), + ("2026-08-17T00:00:00.000Z", "2026-08-17T00:00:00.000Z"), + ("08540", "08540"), + ("012", "012"), + ("08540.5", "08540.5"), + ("0.5", 0.5), + ("10", 10), + ("0", 0), + ("-7", -7), + ("+12", 12), + ("0o17", 15), + ("0x1F", 31), + ("2.50", 2.5), + (".5", 0.5), + ("1e3", 1000.0), + ("~", None), + ("null", None), + ("", None), + ("'08540'", "08540"), + ('"true"', "true"), + ], +) +def test_plain_scalars_follow_deepnote_schema_conventions( + load_yaml: Any, scalar: str, expected: object +) -> None: + assert load_yaml(f"value: {scalar}\n") == {"value": expected} + + +def test_python_object_tags_are_rejected(load_yaml: Any) -> None: + with pytest.raises(yaml.YAMLError): + load_yaml("value: !!python/object/apply:os.getcwd []\n") + + +def test_a_scalar_the_constructor_rejects_raises_a_yaml_error(load_yaml: Any) -> None: + with pytest.raises(yaml.YAMLError): + load_yaml("value: !!int twelve\n") + + +def test_a_repeated_mapping_key_is_rejected(load_yaml: Any) -> None: + with pytest.raises(yaml.YAMLError, match="duplicate key 'notebooks'"): + load_yaml("project:\n notebooks: [a]\n notebooks: [b]\n") + + +def test_the_same_key_may_repeat_in_separate_mappings(load_yaml: Any) -> None: + assert load_yaml("- id: a\n- id: b\n- 1: x\n '1': y\n") == [ + {"id": "a"}, + {"id": "b"}, + {1: "x", "1": "y"}, + ] + + +def test_mapping_tag_on_another_node_is_a_yaml_error(load_yaml: Any) -> None: + with pytest.raises(yaml.YAMLError, match="expected a mapping node"): + load_yaml("!!map [1, 2]") + + +@pytest.mark.parametrize( + "content", ["x: !!timestamp invalid", "x: !!binary [1]", "x: !!bool []"] +) +def test_malformed_explicit_tags_raise_yaml_error(load_yaml, content): + with pytest.raises(yaml.YAMLError): + load_yaml(content) diff --git a/tests/unit/test_streamlit.py b/tests/unit/test_streamlit.py index e3dfb7ec..94531a1a 100644 --- a/tests/unit/test_streamlit.py +++ b/tests/unit/test_streamlit.py @@ -97,3 +97,41 @@ def exists_side_effect(path: str) -> bool: assert mock_logger.warning.call_count == 2 assert mock_venv.start_server.call_count == 1 + + def test_passes_app_id_as_environment_data(self) -> None: + """App IDs are data, never shell syntax; validation belongs to the SDK.""" + apps = [ + { + "id": "11111111-2222-3333-4444-555555555555", + "entrypoint": "a/app.py", + "port": "8501", + }, + {"id": "x; rm -rf /", "entrypoint": "b/app.py", "port": "8502"}, + ] + mock_venv = MagicMock() + + with ( + patch("installer.module.streamlit.fetch_streamlit_apps", return_value=apps), + patch("installer.module.streamlit.os.path.exists", return_value=True), + ): + start_streamlit_servers(mock_venv, MagicMock(spec=logging.Logger)) + + calls = mock_venv.start_server.call_args_list + assert calls[0].args[0].startswith("streamlit run /work/a/app.py ") + assert calls[1].args[0].startswith("streamlit run /work/b/app.py ") + assert calls[0].kwargs["env"] == {"DEEPNOTE_STREAMLIT_APP_ID": apps[0]["id"]} + assert calls[1].kwargs["env"] == {"DEEPNOTE_STREAMLIT_APP_ID": apps[1]["id"]} + + def test_missing_app_id_still_marks_process_as_hosted(self) -> None: + app = {"entrypoint": "app.py", "port": "8501"} + venv = MagicMock() + with ( + patch( + "installer.module.streamlit.fetch_streamlit_apps", return_value=[app] + ), + patch("installer.module.streamlit.os.path.exists", return_value=True), + ): + start_streamlit_servers(venv, MagicMock(spec=logging.Logger)) + assert venv.start_server.call_args.kwargs["env"] == { + "DEEPNOTE_STREAMLIT_APP_ID": "" + } diff --git a/tests/unit/test_streamlit_data_apps.py b/tests/unit/test_streamlit_data_apps.py index f1fc2c69..631bdc22 100644 --- a/tests/unit/test_streamlit_data_apps.py +++ b/tests/unit/test_streamlit_data_apps.py @@ -81,7 +81,7 @@ def test_get_federated_auth_token_raises_when_token_missing(tmp_path, monkeypatc _setup_attached_config(tmp_path, monkeypatch) with patch( - "deepnote_toolkit.streamlit_data_apps._read_streamlit_token_from_context", + "deepnote_toolkit.streamlit_data_apps.read_streamlit_token_from_context", return_value=None, ): with pytest.raises(StreamlitFederatedAuthError) as excinfo: diff --git a/tests/unit/test_virtual_environment.py b/tests/unit/test_virtual_environment.py index 841d0d91..f646495e 100644 --- a/tests/unit/test_virtual_environment.py +++ b/tests/unit/test_virtual_environment.py @@ -98,3 +98,34 @@ def test_import_package_bundle_condition_env_and_priority_mutually_exclusive( condition_env="SOME_ENV_VAR", priority=True, ) + + +def test_server_environment_is_passed_as_data_and_inherits_parent( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Pass app IDs as environment data while retaining the parent environment.""" + import json + import shlex + import sys + + from installer.module.virtual_environment import VirtualEnvironment + + monkeypatch.setenv("TOOLKIT_TEST_PARENT", "inherited") + venv_path = tmp_path / "venv" + (venv_path / "bin").mkdir(parents=True) + (venv_path / "bin" / "activate").write_text("") + result = tmp_path / "result.json" + script = tmp_path / "child.py" + script.write_text( + "import json, os\n" + f"with open({str(result)!r}, 'w') as f:\n" + " json.dump([os.environ['DEEPNOTE_STREAMLIT_APP_ID'], " + "os.environ['TOOLKIT_TEST_PARENT']], f)\n" + ) + app_id = "x; echo must-not-be-executed" + server = VirtualEnvironment(venv_path).start_server( + f"{shlex.quote(sys.executable)} {shlex.quote(str(script))}", + env={"DEEPNOTE_STREAMLIT_APP_ID": app_id}, + ) + assert server.wait(timeout=10) == 0 + assert json.loads(result.read_text()) == [app_id, "inherited"]