diff --git a/tests/fast/source_scan.py b/tests/fast/source_scan.py index 1fbb158afb..7731730b5d 100644 --- a/tests/fast/source_scan.py +++ b/tests/fast/source_scan.py @@ -54,9 +54,10 @@ def imported_modules_of_source(source: str, *, filename: str = "") -> se elif isinstance(node, ast.ImportFrom) and node.module and node.level == 0: imported.add(node.module) elif isinstance(node, ast.Call) and _is_dynamic_import(node.func): + arguments = [*node.args, *(keyword.value for keyword in node.keywords)] imported.update( argument.value - for argument in node.args + for argument in arguments if isinstance(argument, ast.Constant) and isinstance(argument.value, str) ) return imported @@ -66,7 +67,10 @@ def imports_package(modules: Iterable[str], package: str) -> bool: return any(module == package or module.startswith(f"{package}.") for module in modules) +_DYNAMIC_IMPORT_NAMES = ("import_module", "__import__") + + def _is_dynamic_import(func: ast.expr) -> bool: if isinstance(func, ast.Name): - return func.id == "__import__" - return isinstance(func, ast.Attribute) and func.attr in ("import_module", "__import__") + return func.id in _DYNAMIC_IMPORT_NAMES + return isinstance(func, ast.Attribute) and func.attr in _DYNAMIC_IMPORT_NAMES diff --git a/tests/fast/utils/workers/real_ray/test_ray_comm_conformance.py b/tests/fast/utils/workers/real_ray/test_ray_comm_conformance.py new file mode 100644 index 0000000000..2c59906f3e --- /dev/null +++ b/tests/fast/utils/workers/real_ray/test_ray_comm_conformance.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from collections.abc import AsyncIterator, Iterator + +import pytest +import ray +from tests.fast.utils.workers.conformance import ( + POOL_ID, + READY_TIMEOUT_SECONDS, + SHARED_CHECK_IDS, + SHARED_CHECKS, + HandleCheck, + compute_spec, +) +from tests.fast.utils.workers.conftest import worker_manager_args +from tests.fast.utils.workers.real_ray.conftest import ( + kill_named_worker_manager, + kill_quietly, + wait_until_named_manager_is_gone, +) + +from miles.utils.workers.naming import compute_cell_id +from miles.utils.workers.ray_worker_handle import RayWorkerHandle +from miles.utils.workers.ray_worker_manager import RayWorkerManager +from miles.utils.workers.types import WorkerCommBackend +from miles.utils.workers.worker_handle import BaseWorkerHandle +from miles.utils.workers.worker_provider.ray import RayWorkerProvider + +CELL_ID = compute_cell_id(pool_id=POOL_ID, cell_index=0) + +CONFIRM_DEAD_TIMEOUT_SECONDS = 60.0 + + +# every actor this pool launches is a fresh process that imports miles, so the checks below +# share one pool per class rather than paying that twice each +@pytest.fixture(autouse=True, scope="class") +def clean_named_worker_manager(ray_local_mode) -> Iterator[None]: + kill_named_worker_manager() + wait_until_named_manager_is_gone() + yield + kill_named_worker_manager() + wait_until_named_manager_is_gone() + + +@pytest.fixture(scope="class") +def ray_comm_pool(ray_local_mode) -> Iterator[ray.actor.ActorHandle]: + handle = RayWorkerManager.launch( + worker_manager_args(env_report_interval_seconds=0.0), + [compute_spec(rpc_port=0)], + {}, + comm_backend=WorkerCommBackend.RAY, + ) + yield handle + kill_quietly(handle) + + +@pytest.fixture +async def ray_comm_handle(ray_comm_pool: ray.actor.ActorHandle) -> AsyncIterator[BaseWorkerHandle]: + provider = RayWorkerProvider(worker_manager_handle=ray_comm_pool, pool_ids=[POOL_ID]) + handle = provider.get_handle(f"{POOL_ID}-0-0") + await handle.wait_ready(timeout=READY_TIMEOUT_SECONDS) + yield handle + + +class TestARayLaunchedWorkerCalledOverRay: + def test_the_launcher_answers_with_an_actor_handle(self, ray_comm_pool: ray.actor.ActorHandle): + """Both wires stay supported until the default flips, so this column must keep running beside rpc.""" + provider = RayWorkerProvider(worker_manager_handle=ray_comm_pool, pool_ids=[POOL_ID]) + + handle = provider.get_handle(f"{POOL_ID}-0-0") + + assert isinstance(handle, RayWorkerHandle) + + @pytest.mark.parametrize("check", SHARED_CHECKS, ids=SHARED_CHECK_IDS) + async def test_the_handle_contract_holds(self, ray_comm_handle: BaseWorkerHandle, check: HandleCheck): + """The contract a driver is written against must not depend on which wire carries the call.""" + await check(ray_comm_handle) + + +class TestWhenTheLauncherStopsTheCell: + async def test_the_worker_is_confirmed_dead( + self, ray_comm_pool: ray.actor.ActorHandle, ray_comm_handle: BaseWorkerHandle + ): + """Fault tolerance kills a cell and waits for this confirmation before healing it, on either wire.""" + await ray_comm_pool.stop_cells.remote([CELL_ID]) + + await ray_comm_handle.wait_dead(timeout=CONFIRM_DEAD_TIMEOUT_SECONDS) + + assert await ray_comm_handle._probe_is_dead() is True diff --git a/tests/fast/utils/workers/test_ray_communication_boundary.py b/tests/fast/utils/workers/test_ray_communication_boundary.py new file mode 100644 index 0000000000..7fddbc21ce --- /dev/null +++ b/tests/fast/utils/workers/test_ray_communication_boundary.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +import functools +from pathlib import Path + +from tests.fast.source_scan import ( + FRAMEWORK_ROOT, + REPO_ROOT, + imported_modules, + imported_modules_of_source, + imports_package, + relative_paths, + shipped_modules, +) + +EXCLUDED_DIRS = (REPO_ROOT / "tests",) + +RAY_USING_MODULES = { + "miles/ray/placement_group.py": "launcher closure: placement groups are how ray is asked to schedule", + "miles/utils/ray_utils.py": "launcher closure: node lookup and pinning options for the launcher's own calls", + "miles/utils/workers/ray_worker_manager.py": "launcher closure: it is the launcher", + "miles/utils/workers/ray_worker_handle.py": "launcher closure: the handle of the ray communication mode itself", + "miles/utils/workers/worker_provider/ray.py": "launcher closure: it reads the launcher's own bookkeeping", + "miles/utils/workers/backend_capability/ray.py": "launcher closure: it builds the ray provider and operations", + "miles/utils/workers/cell_operations/ray.py": "launcher closure: suspend and resume are launcher verbs", + "miles/utils/external_utils/command_utils/ray_backend/command.py": "launcher closure: running the ray launch scripts on every node the launcher has", + "miles/utils/misc.py": "launcher closure: the node probe every launched actor answers with", + "miles/utils/http_utils.py": "launcher closure: reaching a port on a node the launcher scheduled", + "miles/utils/object_store.py": "object store exemption: the ray object store is a data plane of its own", + "miles/utils/object_store_config.py": "object store exemption: config discovers the local node for its own data plane", + "miles/dashboard/backend.py": "known debt: the dashboard collector is a named actor, tracked outside M23", + "miles/dashboard/collector.py": "known debt: the dashboard collector is a named actor, tracked outside M23", + "miles/dashboard/hooks.py": "known debt: the dashboard reads its gpu ids from ray, tracked outside M23", + "miles/utils/tracking_utils/prometheus_utils.py": "known debt: the prometheus collector is a ray actor, skipped under kubernetes", + "miles/ray/train_actor.py": "launcher closure: a launched actor reads the gpu ids ray gave it", + "miles/backends/fsdp_utils/update_weight_utils.py": "node ip lookup for a collective, not a call to another worker", + "miles/backends/training_utils/weight_update/protocols/broadcast.py": "node ip lookup for a collective, not a call to another worker", + "miles/backends/training_utils/weight_update/protocols/p2p_transfer_utils.py": "node ip lookup for a collective, not a call to another worker", + "miles/utils/debug_utils/replay_reward_fn.py": "tooling: a standalone debugging script", + "miles/utils/test_utils/mock_sglang_engine.py": "tooling: a test double that stands in for a ray-launched engine", + "tools/convert_torch_dist_to_hf_ray.py": "tooling: a standalone conversion script that fans out over a ray cluster", + "examples/experimental/formal_math/single_round/kimina_wrapper.py": "user example: a verifier pool of its own, outside the worker layer", + "examples/experimental/formal_math/single_round/reward_fn.py": "user example: a verifier pool of its own, outside the worker layer", +} + + +@functools.cache +def _scanned_modules() -> tuple[Path, ...]: + return tuple(shipped_modules(exclude_dirs=EXCLUDED_DIRS)) + + +@functools.cache +def _imports_ray(path: Path) -> bool: + return imports_package(imported_modules(path), "ray") + + +@functools.cache +def _ray_using_module_paths() -> tuple[str, ...]: + return tuple(relative_paths(path for path in _scanned_modules() if _imports_ray(path))) + + +class TestRayIsOnlyUsedInsideTheLauncherClosure: + def test_no_module_reaches_for_ray_without_being_listed(self): + """Ray communication outside the launcher closure is what the rpc comm backend exists to remove.""" + unlisted = sorted(set(_ray_using_module_paths()) - set(RAY_USING_MODULES)) + + assert unlisted == [], ( + f"{unlisted} import ray; a worker must be reachable over rpc too, so either drop the import or " + f"add it to RAY_USING_MODULES with the reason it belongs to the launcher closure" + ) + + def test_the_list_names_no_module_that_stopped_using_ray(self): + """A stale exemption reads as permission to bring ray back into a module that no longer needs it.""" + stale = sorted(set(RAY_USING_MODULES) - set(_ray_using_module_paths())) + + assert stale == [] + + def test_every_listed_module_says_why(self): + """The exemptions are a ledger of remaining debt, and a blank reason hides an entry from review.""" + assert sorted(name for name, reason in RAY_USING_MODULES.items() if not reason.strip()) == [] + + def test_the_rpc_layer_itself_never_touches_ray(self): + """The rpc client and server must run in a process that has no ray at all, such as a kubernetes pod.""" + rpc_modules = [path for path in _scanned_modules() if "rpc" in path.parts and _imports_ray(path)] + + assert rpc_modules == [] + + def test_a_served_worker_is_started_without_ray(self): + """serve_actor happens to run inside a ray actor, but it starts the same server a pod starts.""" + assert not _imports_ray(FRAMEWORK_ROOT / "utils" / "workers" / "serving" / "serve_actor.py") + + +class TestTheScanReachesEveryProcessTheDriverIsPartOf: + def test_the_orchestration_scripts_are_scanned(self): + """The driver is where a ray-only exception handler hid, and miles/ alone never covered it.""" + scanned = relative_paths(_scanned_modules()) + + assert {"train.py", "train_async.py", "train_multi_lora_async.py"} <= set(scanned) + + def test_a_driver_script_that_reaches_for_ray_would_be_reported(self): + """A check that only ever looks under miles/ passes on the very file that broke under rpc.""" + assert imports_package(imported_modules_of_source("import ray\n"), "ray") + + +class TestTheShapesThatUsedToSlipThrough: + def test_a_submodule_import_counts(self): + """`import ray.exceptions` reaches ray just as much as `import ray` does.""" + assert imports_package(imported_modules_of_source("import ray.exceptions\n"), "ray") + + def test_an_aliased_import_counts(self): + """Renaming the module on the way in does not rename what it talks to.""" + assert imports_package(imported_modules_of_source("import ray as r\n"), "ray") + + def test_a_dynamic_import_counts(self): + """importlib is the shape an import lands in once someone wants it to not look like one.""" + assert imports_package(imported_modules_of_source('import importlib\nimportlib.import_module("ray")\n'), "ray") + + def test_a_module_that_merely_shares_the_prefix_does_not_count(self): + """A false positive costs the ledger its meaning, and `raydium` is not ray.""" + assert not imports_package(imported_modules_of_source("import raydium\n"), "ray")