mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Fence ray to the launcher and check both wires agree (#2620)
This commit is contained in:
@@ -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")
|
||||
Reference in New Issue
Block a user