diff --git a/examples/create_and_wait_machine.py b/examples/create_and_wait_machine.py new file mode 100644 index 0000000..97c2d1a --- /dev/null +++ b/examples/create_and_wait_machine.py @@ -0,0 +1,42 @@ +"""Create a Dedalus Machine and wait until it is running. + +Usage: + export DEDALUS_API_KEY=... + python examples/create_and_wait_machine.py +""" + +from __future__ import annotations + +import os +import sys + +from dedalus_sdk import Dedalus +from dedalus_sdk.lib.machine_wait import create_and_wait, MachineWaitError + + +def main() -> int: + if not os.environ.get("DEDALUS_API_KEY") and not os.environ.get("DEDALUS_X_API_KEY"): + print("Set DEDALUS_API_KEY (or DEDALUS_X_API_KEY) first.", file=sys.stderr) + return 1 + + client = Dedalus() + + print("Creating machine and waiting until running...") + try: + machine = create_and_wait( + client, + memory_mib=2048, + storage_gib=10, + vcpu=1, + on_status=lambda m: print(f" phase={m.status.phase} reason={m.status.reason}"), + ) + except MachineWaitError as exc: + print(f"Failed: {exc}", file=sys.stderr) + return 1 + + print(f"Ready: {machine.machine_id} ({machine.status.phase})") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/dedalus_sdk/lib/machine_wait.py b/src/dedalus_sdk/lib/machine_wait.py new file mode 100644 index 0000000..9057022 --- /dev/null +++ b/src/dedalus_sdk/lib/machine_wait.py @@ -0,0 +1,180 @@ +"""Machine lifecycle wait helpers. + +These helpers sit on top of the generated SDK and implement the common +"create a machine, then wait until it is usable" workflow that official +examples currently hand-roll with sleep + retrieve loops. + +Example:: + + from dedalus_sdk import Dedalus + from dedalus_sdk.lib.machine_wait import create_and_wait, wait_until_running + + client = Dedalus() + machine = create_and_wait( + client, + memory_mib=2048, + storage_gib=10, + vcpu=1, + on_status=lambda m: print(m.status.phase), + ) +""" + +from __future__ import annotations + +import time +from typing import Any, Callable, Iterable, Optional, Sequence, Union + +from dedalus_sdk import Dedalus +from dedalus_sdk.types.machine import Machine + +MachinePhase = str +OnStatus = Callable[[Machine], None] + +TERMINAL_PHASES = frozenset({"failed", "destroyed"}) + + +class MachineWaitError(Exception): + """Raised when waiting for a machine fails or times out.""" + + +class MachineTerminalError(MachineWaitError): + """Machine reached a terminal phase (failed / destroyed).""" + + def __init__(self, machine_id: str, phase: str, last_error: Optional[str] = None): + self.machine_id = machine_id + self.phase = phase + self.last_error = last_error + msg = f"Machine {machine_id} reached terminal phase '{phase}'" + if last_error: + msg = f"{msg}: {last_error}" + super().__init__(msg) + + +class MachineWaitTimeout(MachineWaitError): + """Timed out while waiting for a machine phase.""" + + def __init__(self, machine_id: str, timeout: float, last_phase: str): + self.machine_id = machine_id + self.timeout = timeout + self.last_phase = last_phase + super().__init__( + f"Timed out waiting for machine {machine_id} after {timeout:.1f}s " + f"(last phase: {last_phase})" + ) + + +def wait_until( + client: Dedalus, + machine_id: str, + *, + predicate: Callable[[Machine], bool], + timeout: float = 120.0, + poll_interval: float = 1.5, + on_status: Optional[OnStatus] = None, +) -> Machine: + """Poll until ``predicate(machine)`` is true. + + Raises: + MachineTerminalError: if phase is ``failed`` or ``destroyed`` + MachineWaitTimeout: if ``timeout`` is exceeded + """ + deadline = time.monotonic() + timeout + + machine = client.machines.retrieve(machine_id=machine_id) + if on_status: + on_status(machine) + if predicate(machine): + return machine + + while True: + phase = getattr(machine.status, "phase", None) or "" + if phase in TERMINAL_PHASES: + last_error = getattr(machine.status, "last_error", None) + raise MachineTerminalError(machine_id, phase, last_error) + + remaining = deadline - time.monotonic() + if remaining <= 0: + raise MachineWaitTimeout(machine_id, timeout, phase) + + time.sleep(min(poll_interval, remaining)) + + machine = client.machines.retrieve(machine_id=machine_id) + if on_status: + on_status(machine) + if predicate(machine): + return machine + + +def wait_until_running( + client: Dedalus, + machine_id: str, + *, + timeout: float = 120.0, + poll_interval: float = 1.5, + on_status: Optional[OnStatus] = None, +) -> Machine: + """Wait until ``status.phase == "running"``.""" + return wait_until( + client, + machine_id, + predicate=lambda m: getattr(m.status, "phase", None) == "running", + timeout=timeout, + poll_interval=poll_interval, + on_status=on_status, + ) + + +def wait_until_phase( + client: Dedalus, + machine_id: str, + phase: Union[str, Sequence[str]], + *, + timeout: float = 120.0, + poll_interval: float = 1.5, + on_status: Optional[OnStatus] = None, +) -> Machine: + """Wait until the machine reaches one of the given phases.""" + phases = {phase} if isinstance(phase, str) else set(phase) + return wait_until( + client, + machine_id, + predicate=lambda m: getattr(m.status, "phase", None) in phases, + timeout=timeout, + poll_interval=poll_interval, + on_status=on_status, + ) + + +def create_and_wait( + client: Dedalus, + *, + memory_mib: int, + storage_gib: int, + vcpu: Union[int, float], + autosleep: Optional[str] = None, + timeout: float = 120.0, + poll_interval: float = 1.5, + on_status: Optional[OnStatus] = None, + **create_kwargs: Any, +) -> Machine: + """Create a machine and wait until it is running. + + Extra keyword args are forwarded to ``client.machines.create``. + """ + params: dict[str, Any] = { + "memory_mib": memory_mib, + "storage_gib": storage_gib, + "vcpu": vcpu, + **create_kwargs, + } + if autosleep is not None: + params["autosleep"] = autosleep + + machine = client.machines.create(**params) + return wait_until_running( + client, + machine.machine_id, + timeout=timeout, + poll_interval=poll_interval, + on_status=on_status, + ) diff --git a/tests/test_machine_wait.py b/tests/test_machine_wait.py new file mode 100644 index 0000000..ae470a3 --- /dev/null +++ b/tests/test_machine_wait.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from dedalus_sdk.lib.machine_wait import ( + MachineTerminalError, + MachineWaitTimeout, + wait_until, + wait_until_running, +) + + +def _machine(phase: str, last_error: str | None = None): + status = SimpleNamespace( + phase=phase, + reason=phase, + last_error=last_error, + ) + return SimpleNamespace(machine_id="dm-test", status=status) + + +def test_wait_until_running_succeeds(): + phases = ["accepted", "starting", "running"] + client = MagicMock() + client.machines.retrieve.side_effect = [_machine(p) for p in phases] + + result = wait_until_running(client, "dm-test", poll_interval=0.01, timeout=2.0) + assert result.status.phase == "running" + assert client.machines.retrieve.call_count == 3 + + +def test_wait_until_terminal_failed(): + client = MagicMock() + client.machines.retrieve.return_value = _machine("failed", last_error="no capacity") + + with pytest.raises(MachineTerminalError) as ei: + wait_until( + client, + "dm-test", + predicate=lambda m: m.status.phase == "running", + poll_interval=0.01, + timeout=1.0, + ) + assert "failed" in str(ei.value) + assert "no capacity" in str(ei.value) + + +def test_wait_until_timeout(): + client = MagicMock() + client.machines.retrieve.return_value = _machine("starting") + + with pytest.raises(MachineWaitTimeout): + wait_until_running(client, "dm-test", poll_interval=0.01, timeout=0.05)