Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_modal_spawn_executor.py411 linesDownload Raw Back to tests
1"""Tests for ModalSpawnExecutor — the v0-finished Modal-backed executor.2 3These tests exercise the executor's contract WITHOUT requiring a live4Modal connection. They use a mock `modal.Function` that records calls5and returns canned `FunctionCall`-shaped objects.6 7For end-to-end Modal integration testing, see the manual ops runbook in8`~/.hermes/scripts/composer-modal/README.md` — that requires a real9Modal account and incurs spend.10"""11from __future__ import annotations12 13import importlib14import time15import pytest16 17from composer_replication.diloco.serverless import (18    ModalSpawnExecutor,19    ReplicaHandle,20)21 22 23def _is_modal_installed() -> bool:24    try:25        importlib.import_module("modal")26        return True27    except ImportError:28        return False29 30 31# ---------------------------------------------------------------------32# Mock infrastructure33# ---------------------------------------------------------------------34 35 36class _MockFunctionCall:37    """Stand-in for `modal.functions.FunctionCall`.38 39    Behavior knobs:40    - `result_value`: what `.get()` returns when called after `delay_s`41    - `delay_s`: how many seconds before `.get()` stops raising TimeoutError42    - `raise_on_get`: if set, `.get()` raises this exception class instead43    - `_creation_time`: monotonic timestamp at construction44    """45    _next_id = 046 47    def __init__(self, result_value=None, *, delay_s=0.0, raise_on_get=None):48        self.object_id = f"fc-mock-{_MockFunctionCall._next_id:04d}"49        _MockFunctionCall._next_id += 150        self._result = result_value51        self._delay_s = delay_s52        self._raise_on_get = raise_on_get53        self._creation_time = time.monotonic()54        self._cancelled = False55 56    def get(self, timeout=None):57        if self._cancelled:58            raise RuntimeError("FunctionCall was cancelled")59        elapsed = time.monotonic() - self._creation_time60        if elapsed < self._delay_s:61            # Not ready yet62            if timeout is None or timeout <= 0:63                raise TimeoutError(f"Mock not ready after {elapsed:.3f}s")64            # Wait up to the timeout for the delay to elapse65            wait = min(timeout, self._delay_s - elapsed)66            time.sleep(wait)67            elapsed = time.monotonic() - self._creation_time68            if elapsed < self._delay_s:69                raise TimeoutError(f"Mock still not ready after wait")70        if self._raise_on_get is not None:71            raise self._raise_on_get72        return self._result73 74    def cancel(self):75        self._cancelled = True76 77    def get_dashboard_url(self):78        return f"https://modal.com/calls/{self.object_id}"79 80 81class _MockModalFunction:82    """Mock of a `@app.function`-decorated callable.83 84    Captures `.spawn()` arg-tuples for assertions. Returns85    `_MockFunctionCall` instances.86    """87 88    def __init__(self, *, fcall_factory=None):89        self.spawn_calls: list[tuple[tuple, dict]] = []90        self._fcall_factory = fcall_factory or (lambda **kw: _MockFunctionCall(91            result_value={"rank": kw.get("rank"), "ok": True},92        ))93        # The duck-type contract: ModalSpawnExecutor checks for .spawn and .remote94        self.app = None  # No deploy needed95 96    def spawn(self, *args, **kwargs):97        self.spawn_calls.append((args, kwargs))98        return self._fcall_factory(**kwargs)99 100    def remote(self, *args, **kwargs):101        # Required by the duck-type check — not exercised in spawn flow102        return self.spawn(*args, **kwargs).get(timeout=60)103 104 105# ---------------------------------------------------------------------106# Construction / preconditions107# ---------------------------------------------------------------------108 109 110@pytest.mark.skipif(not _is_modal_installed(),111                    reason="modal not installed in this venv")112def test_modal_spawn_executor_rejects_non_function():113    with pytest.raises(TypeError, match="modal_function must be"):114        ModalSpawnExecutor(modal_function="not a function")115    with pytest.raises(TypeError, match="modal_function must be"):116        ModalSpawnExecutor(modal_function=lambda x: x)117    with pytest.raises(TypeError, match=".spawn"):118        ModalSpawnExecutor(modal_function=object())119 120 121@pytest.mark.skipif(not _is_modal_installed(),122                    reason="modal not installed in this venv")123def test_modal_spawn_executor_accepts_mock_with_spawn_and_remote():124    mock_fn = _MockModalFunction()125    executor = ModalSpawnExecutor(modal_function=mock_fn)126    assert executor.modal_function is mock_fn127    assert executor.backend_name == "modal_spawn"128    assert executor.supports_inter_replica_network is False129 130 131def test_modal_spawn_executor_missing_modal_raises_runtime_error():132    """If `modal` is genuinely missing, the import-error path should fire.133 134    Only meaningful in venvs without modal — when modal is installed, the135    import-failure path is unreachable without monkey-patching the136    function-local import (brittle across CPython versions). The137    skeleton-executor test in test_skeleton_executors.py covers the138    "modal absent" contract from a different angle and is the canonical139    pin for that path.140    """141    if _is_modal_installed():142        pytest.skip(143            "modal is installed in this venv; the missing-module path "144            "is covered by test_skeleton_executors.py in venvs without modal"145        )146    with pytest.raises(RuntimeError, match="modal client"):147        ModalSpawnExecutor(modal_function=_MockModalFunction())148 149 150# ---------------------------------------------------------------------151# launch_replicas152# ---------------------------------------------------------------------153 154 155@pytest.mark.skipif(not _is_modal_installed(),156                    reason="modal not installed in this venv")157def test_launch_replicas_calls_spawn_n_times_with_rank_kwarg():158    mock_fn = _MockModalFunction()159    executor = ModalSpawnExecutor(modal_function=mock_fn)160 161    handles = executor.launch_replicas(162        n_replicas=4,163        entrypoint="ignored",  # Pinned via decorator164        entrypoint_args={"rendezvous_uri": "/vol/run42", "world_size": 4},165    )166 167    assert len(handles) == 4168    # Each rank in order, backend correct, call_id captured169    for i, h in enumerate(handles):170        assert h.rank == i171        assert h.backend_name == "modal_spawn"172        assert "call_id" in h.metadata173        assert h.metadata["call_id"].startswith("fc-mock-")174 175    # spawn() was called 4× with explicit rank + the user kwargs176    assert len(mock_fn.spawn_calls) == 4177    for i, (args, kwargs) in enumerate(mock_fn.spawn_calls):178        assert args == (), f"args should be empty, got {args}"179        assert kwargs["rank"] == i180        assert kwargs["rendezvous_uri"] == "/vol/run42"181        assert kwargs["world_size"] == 4182 183 184@pytest.mark.skipif(not _is_modal_installed(),185                    reason="modal not installed in this venv")186def test_launch_replicas_strips_rank_env_kwarg():187    """`rank_env` is the LocalProcessExecutor convention; ModalSpawn should drop it."""188    mock_fn = _MockModalFunction()189    executor = ModalSpawnExecutor(modal_function=mock_fn)190 191    executor.launch_replicas(192        n_replicas=2,193        entrypoint="ignored",194        entrypoint_args={"rank_env": "REPLICA_RANK", "rendezvous_uri": "/x"},195    )196 197    for _, kwargs in mock_fn.spawn_calls:198        assert "rank_env" not in kwargs199        assert kwargs["rendezvous_uri"] == "/x"200 201 202@pytest.mark.skipif(not _is_modal_installed(),203                    reason="modal not installed in this venv")204def test_launch_replicas_rejects_zero_or_negative():205    mock_fn = _MockModalFunction()206    executor = ModalSpawnExecutor(modal_function=mock_fn)207    with pytest.raises(ValueError, match="n_replicas"):208        executor.launch_replicas(n_replicas=0, entrypoint="x", entrypoint_args={})209    with pytest.raises(ValueError, match="n_replicas"):210        executor.launch_replicas(n_replicas=-1, entrypoint="x", entrypoint_args={})211 212 213@pytest.mark.skipif(not _is_modal_installed(),214                    reason="modal not installed in this venv")215def test_launch_replicas_cancels_prior_on_partial_failure():216    """If spawn fails at rank 2 of 4, the already-launched 0/1 must be cancelled."""217    cancelled_calls = []218 219    def factory(**kwargs):220        rank = kwargs.get("rank", -1)221        if rank == 2:222            raise RuntimeError("simulated spawn failure at rank 2")223        fc = _MockFunctionCall(result_value={"rank": rank})224        original_cancel = fc.cancel225 226        def tracked_cancel():227            cancelled_calls.append(fc.object_id)228            original_cancel()229        fc.cancel = tracked_cancel230        return fc231 232    mock_fn = _MockModalFunction(fcall_factory=factory)233    executor = ModalSpawnExecutor(modal_function=mock_fn)234 235    with pytest.raises(RuntimeError, match="rank=2"):236        executor.launch_replicas(237            n_replicas=4,238            entrypoint="x",239            entrypoint_args={"world_size": 4},240        )241 242    # Ranks 0 and 1 were spawned and should have been cancelled243    assert len(cancelled_calls) == 2244 245 246# ---------------------------------------------------------------------247# poll / collect248# ---------------------------------------------------------------------249 250 251@pytest.mark.skipif(not _is_modal_installed(),252                    reason="modal not installed in this venv")253def test_poll_returns_succeeded_for_immediate_result():254    def factory(**kwargs):255        return _MockFunctionCall(result_value={"rank": kwargs["rank"], "ok": True})256 257    mock_fn = _MockModalFunction(fcall_factory=factory)258    executor = ModalSpawnExecutor(modal_function=mock_fn)259    handles = executor.launch_replicas(260        n_replicas=2,261        entrypoint="x",262        entrypoint_args={"world_size": 2},263    )264 265    # Immediate get() should succeed266    for h in handles:267        status = executor.poll(h)268        assert status == "succeeded", f"rank {h.rank} expected succeeded, got {status}"269 270 271@pytest.mark.skipif(not _is_modal_installed(),272                    reason="modal not installed in this venv")273def test_poll_returns_running_for_delayed_call():274    def factory(**kwargs):275        return _MockFunctionCall(276            result_value={"rank": kwargs["rank"]}, delay_s=10.0,277        )278 279    mock_fn = _MockModalFunction(fcall_factory=factory)280    executor = ModalSpawnExecutor(modal_function=mock_fn)281    handles = executor.launch_replicas(282        n_replicas=1, entrypoint="x", entrypoint_args={},283    )284 285    status = executor.poll(handles[0])286    assert status == "running"287 288 289@pytest.mark.skipif(not _is_modal_installed(),290                    reason="modal not installed in this venv")291def test_poll_returns_failed_for_user_exception():292    def factory(**kwargs):293        return _MockFunctionCall(294            raise_on_get=ValueError("user code blew up"),295        )296 297    mock_fn = _MockModalFunction(fcall_factory=factory)298    executor = ModalSpawnExecutor(modal_function=mock_fn)299    handles = executor.launch_replicas(300        n_replicas=1, entrypoint="x", entrypoint_args={},301    )302 303    status = executor.poll(handles[0])304    assert status == "failed"305 306 307@pytest.mark.skipif(not _is_modal_installed(),308                    reason="modal not installed in this venv")309def test_collect_returns_per_replica_dicts():310    def factory(**kwargs):311        return _MockFunctionCall(result_value={"rank": kwargs["rank"], "ok": True})312 313    mock_fn = _MockModalFunction(fcall_factory=factory)314    executor = ModalSpawnExecutor(modal_function=mock_fn)315    handles = executor.launch_replicas(316        n_replicas=3, entrypoint="x", entrypoint_args={},317    )318 319    results = executor.collect(handles, timeout=5)320    assert len(results) == 3321    for i, r in enumerate(results):322        assert r["rank"] == i323        assert r["status"] == "succeeded"324        assert r["exit_code"] == 0325        assert r["error"] is None326        assert r["result"] == {"rank": i, "ok": True}327        assert r["call_id"].startswith("fc-mock-")328 329 330@pytest.mark.skipif(not _is_modal_installed(),331                    reason="modal not installed in this venv")332def test_collect_caches_results_and_does_not_call_get_twice():333    """Once a poll() succeeds, collect() must read from cache, not call .get() again."""334    get_calls = []335 336    class _CountingFC(_MockFunctionCall):337        def get(self, timeout=None):338            get_calls.append(self.object_id)339            return super().get(timeout=timeout)340 341    def factory(**kwargs):342        return _CountingFC(result_value={"rank": kwargs["rank"]})343 344    mock_fn = _MockModalFunction(fcall_factory=factory)345    executor = ModalSpawnExecutor(modal_function=mock_fn)346    handles = executor.launch_replicas(347        n_replicas=2, entrypoint="x", entrypoint_args={},348    )349 350    # Poll all to cache351    for h in handles:352        executor.poll(h)353    n_polls = len(get_calls)354    assert n_polls == 2  # one .get per poll355 356    # Collect should NOT call .get() again357    results = executor.collect(handles, timeout=5)358    assert len(get_calls) == n_polls  # no additional .get() calls359    assert all(r["status"] == "succeeded" for r in results)360 361 362# ---------------------------------------------------------------------363# Logs / cancel364# ---------------------------------------------------------------------365 366 367@pytest.mark.skipif(not _is_modal_installed(),368                    reason="modal not installed in this venv")369def test_stream_logs_includes_dashboard_url_and_call_id():370    mock_fn = _MockModalFunction()371    executor = ModalSpawnExecutor(modal_function=mock_fn)372    handles = executor.launch_replicas(373        n_replicas=1, entrypoint="x", entrypoint_args={},374    )375 376    log_text = executor.stream_logs(handles[0])377    assert "fc-mock-" in log_text378    assert "https://modal.com/" in log_text379 380 381@pytest.mark.skipif(not _is_modal_installed(),382                    reason="modal not installed in this venv")383def test_cancel_calls_fcall_cancel():384    mock_fn = _MockModalFunction()385    executor = ModalSpawnExecutor(modal_function=mock_fn)386    handles = executor.launch_replicas(387        n_replicas=2, entrypoint="x", entrypoint_args={},388    )389 390    fc0 = executor._handles[0]["fcall"]391    assert fc0._cancelled is False392    executor.cancel(handles[0])393    assert fc0._cancelled is True394    # Cancelling rank 1 doesn't affect rank 0 (already cancelled — no-op)395    executor.cancel(handles[1])396    fc1 = executor._handles[1]["fcall"]397    assert fc1._cancelled is True398 399 400@pytest.mark.skipif(not _is_modal_installed(),401                    reason="modal not installed in this venv")402def test_cancel_unknown_handle_is_noop():403    mock_fn = _MockModalFunction()404    executor = ModalSpawnExecutor(modal_function=mock_fn)405    # No replicas launched — handle has rank that doesn't exist in _handles406    fake_handle = ReplicaHandle(407        rank=99, backend_name="modal_spawn", metadata={"call_id": "nonexistent"},408    )409    # Must not raise410    executor.cancel(fake_handle)411