Codeseys/composer-replication-framework
0
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 