diff --git a/nemo_rl/distributed/virtual_cluster.py b/nemo_rl/distributed/virtual_cluster.py index 1e40b0ff5cc..587f53338eb 100644 --- a/nemo_rl/distributed/virtual_cluster.py +++ b/nemo_rl/distributed/virtual_cluster.py @@ -113,12 +113,12 @@ class PY_EXECUTABLES: DEFAULT_SGLANG_ROUTER_PORT_RANGE_HIGH = 8799 DEFAULT_SGLANG_PROMETHEUS_PORT_RANGE_LOW = 8800 DEFAULT_SGLANG_PROMETHEUS_PORT_RANGE_HIGH = 8999 -# vLLM Router control-plane ports occupy the reserved gap between Ray's +# Inference Router control-plane ports occupy the reserved gap between Ray's # management ports and the master / TCPStore range. -DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW = 1320 -DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH = 1360 -DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW = 1360 -DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH = 1400 +DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW = 1320 +DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH = 1360 +DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW = 1360 +DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH = 1400 # Master address / TCPStore range, tucked below the Ray worker-gRPC band (2000+). DEFAULT_MASTER_PORT_RANGE_LOW = 1400 DEFAULT_MASTER_PORT_RANGE_HIGH = 1999 diff --git a/nemo_rl/environments/vllm_router.py b/nemo_rl/environments/inference_router.py similarity index 55% rename from nemo_rl/environments/vllm_router.py rename to nemo_rl/environments/inference_router.py index 2a50c9a1b60..41b60c354e5 100644 --- a/nemo_rl/environments/vllm_router.py +++ b/nemo_rl/environments/inference_router.py @@ -15,25 +15,76 @@ import subprocess import sys import time +from typing import Literal from urllib.error import URLError from urllib.request import urlopen -from pydantic import BaseModel - - -class VllmRouterConfig(BaseModel, extra="forbid"): +from pydantic import BaseModel, model_validator + +RouterBackend = Literal["vllm_router", "smg"] +RouterPolicy = Literal[ + "random", + "round_robin", + "cache_aware", + "power_of_two", + "consistent_hash", + "rendezvous_hash", + "passthrough", + "least_load", + "manual", + "prefix_hash", +] + +_BACKEND_POLICIES = { + "vllm_router": frozenset( + { + "random", + "round_robin", + "cache_aware", + "power_of_two", + "consistent_hash", + } + ), + "smg": frozenset( + { + "random", + "round_robin", + "cache_aware", + "power_of_two", + "consistent_hash", + "passthrough", + "least_load", + "manual", + "prefix_hash", + } + ), +} + + +class InferenceRouterConfig(BaseModel, extra="forbid"): enabled: bool = False - policy: str = "consistent_hash" + backend: RouterBackend = "vllm_router" + policy: RouterPolicy = "consistent_hash" + @model_validator(mode="after") + def validate_backend_policy(self) -> "InferenceRouterConfig": + if self.policy not in _BACKEND_POLICIES[self.backend]: + supported = ", ".join(sorted(_BACKEND_POLICIES[self.backend])) + raise ValueError( + f"Router backend {self.backend!r} does not support policy " + f"{self.policy!r}; supported policies: {supported}" + ) + return self -class VllmRouterProcess: + +class InferenceRouterProcess: def __init__( self, worker_base_urls: list[str], host: str, port: int, prometheus_port: int, - config: VllmRouterConfig, + config: InferenceRouterConfig, ) -> None: self.worker_base_urls = [ base_url.rstrip("/").removesuffix("/v1") for base_url in worker_base_urls @@ -44,16 +95,37 @@ def __init__( self.config = config self._process: subprocess.Popen | None = None + @property + def name(self) -> str: + return "vLLM Router" if self.config.backend == "vllm_router" else "SMG" + + @property + def session_affinity_header(self) -> str: + if self.config.backend == "vllm_router": + return "X-Session-ID" + return "X-SMG-Routing-Key" + @property def command(self) -> list[str]: + if self.config.backend == "vllm_router": + module = "vllm_router.launch_router" + policy = self.config.policy + else: + module = "smg.launch_router" + policy = ( + "consistent_hashing" + if self.config.policy == "consistent_hash" + else self.config.policy + ) + return [ sys.executable, "-m", - "vllm_router.launch_router", + module, "--worker-urls", *self.worker_base_urls, "--policy", - self.config.policy, + policy, "--host", self.host, "--port", @@ -72,7 +144,7 @@ def readiness_url(self) -> str: def start(self) -> None: if self._process is not None: - raise RuntimeError("vLLM Router process has already been started") + raise RuntimeError(f"{self.name} process has already been started") self._process = subprocess.Popen(self.command) def wait_until_ready( @@ -82,7 +154,7 @@ def wait_until_ready( ) -> None: process = self._process if process is None: - raise RuntimeError("vLLM Router process has not been started") + raise RuntimeError(f"{self.name} process has not been started") deadline = time.monotonic() + timeout while True: @@ -90,7 +162,7 @@ def wait_until_ready( if return_code is not None: self._process = None raise RuntimeError( - f"vLLM Router process exited with code {return_code} " + f"{self.name} process exited with code {return_code} " "before becoming ready" ) @@ -106,12 +178,12 @@ def wait_until_ready( if time.monotonic() >= deadline: raise TimeoutError( - f"vLLM Router did not become ready within {timeout} seconds" + f"{self.name} did not become ready within {timeout} seconds" ) time.sleep(poll_interval) - def stop(self, timeout: float = 5.0) -> None: + def stop(self, timeout: float = 10.0) -> None: process = self._process if process is None: return diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index 64d3ea72599..451c04c32c4 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -38,15 +38,18 @@ from nemo_rl.distributed.virtual_cluster import ( DEFAULT_GYM_PORT_RANGE_HIGH, DEFAULT_GYM_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH, - DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW, _get_free_port_local, _get_node_ip_local, ) +from nemo_rl.environments.inference_router import ( + InferenceRouterConfig, + InferenceRouterProcess, +) from nemo_rl.environments.interfaces import EnvironmentInterface -from nemo_rl.environments.vllm_router import VllmRouterConfig, VllmRouterProcess from nemo_rl.models.policy import TokenizerConfig from nemo_rl.utils.routed_experts_codec import decode_routed_experts from nemo_rl.utils.timer import Timer @@ -107,7 +110,7 @@ class NemoGymConfig(TypedDict): model_name: str base_urls: List[str] initial_global_config_dict: Dict[str, Any] - vllm_router: NotRequired[VllmRouterConfig] + router: NotRequired[InferenceRouterConfig] # Port range for Gym HTTP servers (head server + subprocess servers). # Defaults to DEFAULT_GYM_PORT_RANGE_LOW/HIGH (5000-5999) from # nemo_rl.distributed.virtual_cluster. See the port layout there. @@ -361,7 +364,7 @@ class NemoGym(EnvironmentInterface): def __init__(self, cfg: NemoGymConfig): self.cfg = cfg - self._vllm_router: VllmRouterProcess | None = None + self._router: InferenceRouterProcess | None = None # Reconstruct the processor inside the actor (rather than serializing it # per rollout call) for full-trajectory multimodal postprocessing. self._processor: Optional[Any] = None @@ -394,24 +397,24 @@ def _spinup(self) -> None: self.head_server_port = _get_free_port_local(_gym_port_low, _gym_port_high) policy_base_urls = self.cfg["base_urls"] - vllm_router_config = self.cfg.get("vllm_router") - if vllm_router_config is not None and vllm_router_config.enabled: + router_config = self.cfg.get("router") + if router_config is not None and router_config.enabled: router_port = _get_free_port_local( - DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH, ) prometheus_port = _get_free_port_local( - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, ) - self._vllm_router = VllmRouterProcess( + self._router = InferenceRouterProcess( worker_base_urls=self.cfg["base_urls"], host=self.node_ip, port=router_port, prometheus_port=prometheus_port, - config=vllm_router_config, + config=router_config, ) - policy_base_urls = [self._vllm_router.openai_base_url] + policy_base_urls = [self._router.openai_base_url] from nemo_gym.cli import GlobalConfigDictParserConfig, RunHelper from nemo_gym.rollout_collection import RolloutCollectionHelper @@ -426,7 +429,7 @@ def _spinup(self) -> None: initial_global_config_dict = DictConfig( self.cfg.get("initial_global_config_dict") or {} ) - if self._vllm_router is not None: + if self._router is not None: initial_global_config_dict = cast( DictConfig, OmegaConf.merge( @@ -435,7 +438,7 @@ def _spinup(self) -> None: "policy_model": { "responses_api_models": { "vllm_model": { - "session_affinity_header": "X-Session-ID" + "session_affinity_header": self._router.session_affinity_header } } } @@ -497,9 +500,9 @@ def _spinup(self) -> None: self.rh = RunHelper() try: - if self._vllm_router is not None: - self._vllm_router.start() - self._vllm_router.wait_until_ready() + if self._router is not None: + self._router.start() + self._router.wait_until_ready() self.rh.start( global_config_dict_parser_config=GlobalConfigDictParserConfig( dotenv_path=Path(__file__.removesuffix(RELATIVE_PATH)).absolute() @@ -509,8 +512,8 @@ def _spinup(self) -> None: ) ) except Exception: - router = self._vllm_router - self._vllm_router = None + router = self._router + self._router = None if router is not None: router.stop() raise @@ -871,8 +874,8 @@ def shutdown(self) -> None: try: self.rh.shutdown() finally: - router = self._vllm_router - self._vllm_router = None + router = self._router + self._router = None if router is not None: router.stop() @@ -1010,11 +1013,11 @@ def spinup_nemo_gym_actor( The spun-up NemoGym Ray actor handle (_spinup already awaited). """ nemo_gym_dict = dict(env_configs["nemo_gym"]) - vllm_router_dict = nemo_gym_dict.pop("vllm_router", None) - vllm_router_config = ( - VllmRouterConfig.model_validate(vllm_router_dict) - if vllm_router_dict is not None - else VllmRouterConfig() + router_dict = nemo_gym_dict.pop("router", None) + router_config = ( + InferenceRouterConfig.model_validate(router_dict) + if router_dict is not None + else InferenceRouterConfig() ) # NeMo-RL-side detection knobs are top-level NemoGymConfig fields @@ -1035,7 +1038,7 @@ def spinup_nemo_gym_actor( nemo_gym_cfg = NemoGymConfig( model_name=model_name, base_urls=base_urls, - vllm_router=vllm_router_config, + router=router_config, invalid_tool_call_patterns=invalid_tool_call_patterns, thinking_tags=thinking_tags, tokenizer_config=tokenizer_config, diff --git a/pyproject.toml b/pyproject.toml index 7cb1b8cbbc6..4aa54837af9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -213,7 +213,7 @@ nvrx = [ nemo_gym = ["nemo_gym"] [dependency-groups] -nemo_gym_router = ["vllm-router==0.1.15"] +nemo_gym_router = ["smg==1.9.0", "vllm-router==0.1.15"] # This is a default group so that we install these even with bare `uv sync` build = [ diff --git a/pyrefly.toml b/pyrefly.toml index ffcbfb4f077..3711df57641 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -139,7 +139,7 @@ project-includes = [ "nemo_rl/environments/metrics.py", "nemo_rl/environments/rewards.py", "nemo_rl/environments/utils.py", - "nemo_rl/environments/vllm_router.py", + "nemo_rl/environments/inference_router.py", "nemo_rl/environments/vlm_environment.py", "nemo_rl/evals/__init__.py", "nemo_rl/evals/answer_parsing.py", diff --git a/tests/unit/distributed/test_virtual_cluster.py b/tests/unit/distributed/test_virtual_cluster.py index d25bf4184c4..3e3a9f03799 100644 --- a/tests/unit/distributed/test_virtual_cluster.py +++ b/tests/unit/distributed/test_virtual_cluster.py @@ -25,6 +25,10 @@ DEFAULT_GENERATION_PORT_RANGE_LOW, DEFAULT_GYM_PORT_RANGE_HIGH, DEFAULT_GYM_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW, DEFAULT_MASTER_PORT_RANGE_HIGH, DEFAULT_MASTER_PORT_RANGE_LOW, DEFAULT_SGLANG_PROMETHEUS_PORT_RANGE_HIGH, @@ -33,10 +37,6 @@ DEFAULT_SGLANG_ROUTER_PORT_RANGE_LOW, DEFAULT_VLLM_PORT_RANGE_LOW, DEFAULT_VLLM_PORTS_PER_ENGINE, - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH, - DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW, PY_EXECUTABLES, RayVirtualCluster, ResourceInsufficientError, @@ -593,20 +593,21 @@ def test_default_port_ranges_ordered_and_below_ephemeral_floor(): assert DEFAULT_MASTER_PORT_RANGE_LOW > 1024 -def test_vllm_router_default_port_ranges_use_reserved_gap(): +def test_inference_router_default_port_ranges_use_reserved_gap(): assert ( - DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH, ) == (1320, 1360) assert ( - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, ) == (1360, 1400) - assert 1312 < DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW + assert 1312 < DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW assert ( - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH - <= DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH + <= DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW ) assert ( - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH <= DEFAULT_MASTER_PORT_RANGE_LOW + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH + <= DEFAULT_MASTER_PORT_RANGE_LOW ) diff --git a/tests/unit/environments/test_vllm_router.py b/tests/unit/environments/test_inference_router.py similarity index 62% rename from tests/unit/environments/test_vllm_router.py rename to tests/unit/environments/test_inference_router.py index 988f7413886..9e207f3bec5 100644 --- a/tests/unit/environments/test_vllm_router.py +++ b/tests/unit/environments/test_inference_router.py @@ -19,14 +19,14 @@ import pytest -from nemo_rl.environments.vllm_router import ( - VllmRouterConfig, - VllmRouterProcess, +from nemo_rl.environments.inference_router import ( + InferenceRouterConfig, + InferenceRouterProcess, ) -def test_builds_static_router_command_and_openai_base_url() -> None: - router = VllmRouterProcess( +def test_builds_vllm_router_command_and_openai_base_url() -> None: + router = InferenceRouterProcess( worker_base_urls=[ "http://worker-0:8000/v1", "http://worker-1:8001/v1/", @@ -34,7 +34,7 @@ def test_builds_static_router_command_and_openai_base_url() -> None: host="10.0.0.5", port=6100, prometheus_port=6600, - config=VllmRouterConfig(enabled=True), + config=InferenceRouterConfig(enabled=True), ) assert router.command == [ @@ -55,41 +55,86 @@ def test_builds_static_router_command_and_openai_base_url() -> None: ] assert router.openai_base_url == "http://10.0.0.5:6100/v1" assert router.readiness_url == "http://10.0.0.5:6100/readiness" + assert router.session_affinity_header == "X-Session-ID" + + +def test_builds_smg_command_and_maps_session_affinity() -> None: + router = InferenceRouterProcess( + worker_base_urls=["http://worker-0:8000/v1"], + host="10.0.0.5", + port=6100, + prometheus_port=6600, + config=InferenceRouterConfig( + enabled=True, + backend="smg", + policy="consistent_hash", + ), + ) + + assert router.command == [ + sys.executable, + "-m", + "smg.launch_router", + "--worker-urls", + "http://worker-0:8000", + "--policy", + "consistent_hashing", + "--host", + "10.0.0.5", + "--port", + "6100", + "--prometheus-port", + "6600", + ] + assert router.session_affinity_header == "X-SMG-Routing-Key" + + +@pytest.mark.parametrize( + ("backend", "policy"), + [ + ("vllm_router", "least_load"), + ("vllm_router", "rendezvous_hash"), + ("smg", "rendezvous_hash"), + ], +) +def test_rejects_policy_not_supported_by_backend(backend: str, policy: str) -> None: + with pytest.raises(ValueError, match="does not support policy"): + InferenceRouterConfig(backend=backend, policy=policy) def test_starts_and_stops_owned_router_process() -> None: - router = VllmRouterProcess( + router = InferenceRouterProcess( worker_base_urls=["http://worker-0:8000/v1"], host="10.0.0.5", port=6100, prometheus_port=6600, - config=VllmRouterConfig(enabled=True), + config=InferenceRouterConfig(enabled=True), ) process = MagicMock() process.poll.return_value = None with patch( - "nemo_rl.environments.vllm_router.subprocess.Popen", + "nemo_rl.environments.inference_router.subprocess.Popen", return_value=process, ) as popen: router.start() popen.assert_called_once_with(router.command) - router.stop(timeout=2.0) - router.stop(timeout=2.0) + router.stop() + router.stop() process.terminate.assert_called_once_with() - process.wait.assert_called_once_with(timeout=2.0) + process.wait.assert_called_once_with(timeout=10.0) process.kill.assert_not_called() def test_force_kills_router_process_when_shutdown_times_out() -> None: - router = VllmRouterProcess( + router = InferenceRouterProcess( worker_base_urls=["http://worker-0:8000/v1"], host="10.0.0.5", port=6100, prometheus_port=6600, - config=VllmRouterConfig(enabled=True), + config=InferenceRouterConfig(enabled=True), ) process = MagicMock() process.poll.return_value = None @@ -99,7 +144,7 @@ def test_force_kills_router_process_when_shutdown_times_out() -> None: ] with patch( - "nemo_rl.environments.vllm_router.subprocess.Popen", + "nemo_rl.environments.inference_router.subprocess.Popen", return_value=process, ): router.start() @@ -114,12 +159,12 @@ def test_force_kills_router_process_when_shutdown_times_out() -> None: def test_waits_until_router_is_ready() -> None: - router = VllmRouterProcess( + router = InferenceRouterProcess( worker_base_urls=["http://worker-0:8000/v1"], host="10.0.0.5", port=6100, prometheus_port=6600, - config=VllmRouterConfig(enabled=True), + config=InferenceRouterConfig(enabled=True), ) process = MagicMock() process.poll.return_value = None @@ -130,14 +175,14 @@ def test_waits_until_router_is_ready() -> None: with ( patch( - "nemo_rl.environments.vllm_router.subprocess.Popen", + "nemo_rl.environments.inference_router.subprocess.Popen", return_value=process, ), patch( - "nemo_rl.environments.vllm_router.urlopen", + "nemo_rl.environments.inference_router.urlopen", side_effect=[URLError("not ready"), ready_response], ) as urlopen, - patch("nemo_rl.environments.vllm_router.time.sleep") as sleep, + patch("nemo_rl.environments.inference_router.time.sleep") as sleep, ): router.start() router.wait_until_ready(timeout=10.0, poll_interval=0.25) @@ -150,22 +195,22 @@ def test_waits_until_router_is_ready() -> None: def test_readiness_fails_when_router_process_exits() -> None: - router = VllmRouterProcess( + router = InferenceRouterProcess( worker_base_urls=["http://worker-0:8000/v1"], host="10.0.0.5", port=6100, prometheus_port=6600, - config=VllmRouterConfig(enabled=True), + config=InferenceRouterConfig(enabled=True), ) process = MagicMock() process.poll.return_value = 17 with ( patch( - "nemo_rl.environments.vllm_router.subprocess.Popen", + "nemo_rl.environments.inference_router.subprocess.Popen", return_value=process, ), - patch("nemo_rl.environments.vllm_router.urlopen") as urlopen, + patch("nemo_rl.environments.inference_router.urlopen") as urlopen, ): router.start() with pytest.raises(RuntimeError, match="exited with code 17"): @@ -175,26 +220,26 @@ def test_readiness_fails_when_router_process_exits() -> None: def test_readiness_times_out() -> None: - router = VllmRouterProcess( + router = InferenceRouterProcess( worker_base_urls=["http://worker-0:8000/v1"], host="10.0.0.5", port=6100, prometheus_port=6600, - config=VllmRouterConfig(enabled=True), + config=InferenceRouterConfig(enabled=True), ) process = MagicMock() process.poll.return_value = None with ( patch( - "nemo_rl.environments.vllm_router.subprocess.Popen", + "nemo_rl.environments.inference_router.subprocess.Popen", return_value=process, ), patch( - "nemo_rl.environments.vllm_router.urlopen", + "nemo_rl.environments.inference_router.urlopen", side_effect=URLError("not ready"), ), - patch("nemo_rl.environments.vllm_router.time.sleep") as sleep, + patch("nemo_rl.environments.inference_router.time.sleep") as sleep, ): router.start() with pytest.raises(TimeoutError, match="did not become ready"): diff --git a/tests/unit/environments/test_nemo_gym.py b/tests/unit/environments/test_nemo_gym.py index 0b7753a11a7..a4df4ca8409 100644 --- a/tests/unit/environments/test_nemo_gym.py +++ b/tests/unit/environments/test_nemo_gym.py @@ -32,10 +32,14 @@ from nemo_rl.distributed.virtual_cluster import ( DEFAULT_GYM_PORT_RANGE_HIGH, DEFAULT_GYM_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH, - DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW, +) +from nemo_rl.environments.inference_router import ( + InferenceRouterConfig, + RouterBackend, ) from nemo_rl.environments.nemo_gym import ( NemoGym, @@ -46,7 +50,6 @@ spinup_nemo_gym_actor, validate_reward_components_match_scalar, ) -from nemo_rl.environments.vllm_router import VllmRouterConfig from nemo_rl.experience.rollouts import _reattach_original_multimodal_payloads from nemo_rl.models.generation.vllm import VllmGeneration @@ -154,7 +157,7 @@ def test_nemo_gym_stub_module(): ) -def test_spinup_nemo_gym_actor_extracts_vllm_router_config(): +def test_spinup_nemo_gym_actor_extracts_router_config(): actor = MagicMock() spinup_ref = object() actor._spinup.remote.return_value = spinup_ref @@ -172,8 +175,9 @@ def test_spinup_nemo_gym_actor_extracts_vllm_router_config(): result = spinup_nemo_gym_actor( env_configs={ "nemo_gym": { - "vllm_router": { + "router": { "enabled": True, + "backend": "smg", "policy": "consistent_hash", }, "config_paths": ["responses_api_models/vllm_model/config.yaml"], @@ -187,25 +191,41 @@ def test_spinup_nemo_gym_actor_extracts_vllm_router_config(): ) actor_cfg = nemo_gym_actor.options.return_value.remote.call_args.args[0] - assert actor_cfg["vllm_router"] == VllmRouterConfig( + assert actor_cfg["router"] == InferenceRouterConfig( enabled=True, + backend="smg", policy="consistent_hash", ) - assert "vllm_router" not in actor_cfg["initial_global_config_dict"] + assert "router" not in actor_cfg["initial_global_config_dict"] ray_get.assert_called_once_with(spinup_ref) assert result is actor -def test_nemo_gym_spinup_routes_policy_requests_through_vllm_router(): +@pytest.mark.parametrize( + ("backend", "session_affinity_header"), + [ + ("vllm_router", "X-Session-ID"), + ("smg", "X-SMG-Routing-Key"), + ], +) +def test_nemo_gym_spinup_routes_policy_requests_through_router( + backend: RouterBackend, + session_affinity_header: str, +) -> None: events = [] router = MagicMock() router.openai_base_url = "http://10.0.0.5:1325/v1" + router.session_affinity_header = session_affinity_header router.start.side_effect = lambda: events.append("router.start") router.wait_until_ready.side_effect = lambda: events.append("router.ready") run_helper = MagicMock() run_helper.start.side_effect = lambda **_: events.append("gym.start") - router_config = VllmRouterConfig(enabled=True, policy="consistent_hash") + router_config = InferenceRouterConfig( + enabled=True, + backend=backend, + policy="consistent_hash", + ) gym = MagicMock() gym.cfg = NemoGymConfig( model_name="test-model", @@ -220,7 +240,7 @@ def test_nemo_gym_spinup_routes_policy_requests_through_vllm_router(): } } }, - vllm_router=router_config, + router=router_config, ) with ( @@ -233,9 +253,8 @@ def test_nemo_gym_spinup_routes_policy_requests_through_vllm_router(): side_effect=[5500, 1325, 1365], ) as get_free_port, patch( - "nemo_rl.environments.nemo_gym.VllmRouterProcess", + "nemo_rl.environments.nemo_gym.InferenceRouterProcess", return_value=router, - create=True, ) as router_process, patch("nemo_gym.cli.RunHelper", return_value=run_helper), patch("nemo_gym.cli.GlobalConfigDictParserConfig") as parser_config, @@ -252,12 +271,12 @@ def test_nemo_gym_spinup_routes_policy_requests_through_vllm_router(): assert get_free_port.call_args_list == [ call(DEFAULT_GYM_PORT_RANGE_LOW, DEFAULT_GYM_PORT_RANGE_HIGH), call( - DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH, ), call( - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW, - DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW, + DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH, ), ] router_process.assert_called_once_with( @@ -274,14 +293,15 @@ def test_nemo_gym_spinup_routes_policy_requests_through_vllm_router(): vllm_model_config = global_config["policy_model"]["responses_api_models"][ "vllm_model" ] - assert vllm_model_config["session_affinity_header"] == "X-Session-ID" + assert vllm_model_config["session_affinity_header"] == session_affinity_header assert vllm_model_config["num_workers"] == 4 @pytest.mark.parametrize("failure_point", ["router.ready", "gym.start"]) -def test_nemo_gym_spinup_stops_vllm_router_on_failure(failure_point): +def test_nemo_gym_spinup_stops_router_on_failure(failure_point): router = MagicMock() router.openai_base_url = "http://10.0.0.5:1325/v1" + router.session_affinity_header = "X-Session-ID" run_helper = MagicMock() failure = RuntimeError(failure_point) if failure_point == "router.ready": @@ -294,7 +314,7 @@ def test_nemo_gym_spinup_stops_vllm_router_on_failure(failure_point): model_name="test-model", base_urls=["http://worker-0:8000/v1"], initial_global_config_dict={}, - vllm_router=VllmRouterConfig(enabled=True), + router=InferenceRouterConfig(enabled=True), ) with ( @@ -307,7 +327,7 @@ def test_nemo_gym_spinup_stops_vllm_router_on_failure(failure_point): side_effect=[5500, 1325, 1365], ), patch( - "nemo_rl.environments.nemo_gym.VllmRouterProcess", + "nemo_rl.environments.nemo_gym.InferenceRouterProcess", return_value=router, ), patch("nemo_gym.cli.RunHelper", return_value=run_helper), @@ -326,12 +346,12 @@ def test_nemo_gym_spinup_stops_vllm_router_on_failure(failure_point): router.stop.assert_called_once_with() -def test_nemo_gym_shutdown_stops_vllm_router_after_gym_on_failure(): +def test_nemo_gym_shutdown_stops_router_after_gym_on_failure(): events = [] router = MagicMock() router.stop.side_effect = lambda: events.append("router.stop") gym = MagicMock() - gym._vllm_router = router + gym._router = router def fail_gym_shutdown(): events.append("gym.shutdown") @@ -343,7 +363,7 @@ def fail_gym_shutdown(): NemoGym.__ray_metadata__.modified_class.shutdown(gym) assert events == ["gym.shutdown", "router.stop"] - assert gym._vllm_router is None + assert gym._router is None @pytest.fixture(scope="function") diff --git a/uv.lock b/uv.lock index 14684d93c57..e750e03808c 100644 --- a/uv.lock +++ b/uv.lock @@ -4439,6 +4439,7 @@ docs = [ { name = "swagger-plugin-for-sphinx" }, ] nemo-gym-router = [ + { name = "smg" }, { name = "vllm-router" }, ] test = [ @@ -4591,7 +4592,10 @@ docs = [ { name = "sphinxcontrib-mermaid" }, { name = "swagger-plugin-for-sphinx" }, ] -nemo-gym-router = [{ name = "vllm-router", specifier = "==0.1.15" }] +nemo-gym-router = [ + { name = "smg", specifier = "==1.9.0" }, + { name = "vllm-router", specifier = "==0.1.15" }, +] test = [ { name = "pytest", specifier = ">=8.4.2" }, { name = "pytest-asyncio" }, @@ -7013,6 +7017,24 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/c3/96/325b8c507ccecc50421fecc0345a502ee6e4a44785af3c4e6ecbadad624a/smart_open-8.0.1-py3-none-any.whl", hash = "sha256:3e97f90e92a952cb57863dfe132082c400a52eeeb27c067692fb51dbcc5b0089", size = 73504, upload-time = "2026-07-15T13:56:09.033Z" }, ] +[[package]] +name = "smg" +version = "1.9.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "grpcio" }, + { name = "grpcio-health-checking" }, + { name = "pyyaml" }, + { name = "setproctitle" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/13/c8/7707780338e7bd79769ba521a86f73c8272468c58006685ab9b2feacf221/smg-1.9.0.tar.gz", hash = "sha256:a075234d6d21a2aea0494ad599bfce2a81d11e8a1a1d9bb4e4f0f12f5a2031b6", size = 3163976, upload-time = "2026-07-30T06:29:14.463Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/3c/2d/1648875495491340ae4353c3c05c04b07b514636bb0d7b408b6120756d32/smg-1.9.0-cp38-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:38261c1de77f10f8b16fffe20d19235f6c888d2e257029f6c0e66071e62d1446", size = 30724795, upload-time = "2026-07-30T06:28:59.565Z" }, + { url = "https://files.pythonhosted.org/packages/e2/ae/9eb904c003c6d45d34f75d97ccd1a3fcb5520c09ee83718e83c013858afe/smg-1.9.0-cp38-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:cb0462657ecfb8179f3781023d61a5ca92ff26f48557b196e5e4f74e5cbad25d", size = 32326605, upload-time = "2026-07-30T06:29:02.879Z" }, + { url = "https://files.pythonhosted.org/packages/9d/a5/2ca1b5bbc9bf8ebbbd23eb77cf7dba30214715693e80c7839502271544ba/smg-1.9.0-cp38-abi3-musllinux_1_1_aarch64.whl", hash = "sha256:3cce9c414271dea3aff43503251176d46d1ca793dcad0eed40aa65e1c7d759f9", size = 32436109, upload-time = "2026-07-30T06:29:05.944Z" }, + { url = "https://files.pythonhosted.org/packages/5c/59/b8d1e3abdfd9e4f901bf2e444d69bf4e914261518090d64f199c3d3e347f/smg-1.9.0-cp38-abi3-musllinux_1_1_x86_64.whl", hash = "sha256:11c2bea9079a15681c5d2c23ca4fdac47d4781a80d85317d1e7f479ab059a9dd", size = 30833940, upload-time = "2026-07-30T06:29:09.243Z" }, +] + [[package]] name = "smg-grpc-proto" version = "0.4.14"