Fence ray to the launcher and check both wires agree (#2620)

This commit is contained in:
fzyzcjy
2026-09-26 18:55:25 +08:00
committed by GitHub
parent ff1b8a9012
commit f21aac53cc
3 changed files with 216 additions and 3 deletions
+7 -3
View File
@@ -54,9 +54,10 @@ def imported_modules_of_source(source: str, *, filename: str = "<source>") -> 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
@@ -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
@@ -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")