diff --git a/python/sglang/srt/arg_groups/serving_hook.py b/python/sglang/srt/arg_groups/serving_hook.py index 404509bb3..2015ebf0e 100644 --- a/python/sglang/srt/arg_groups/serving_hook.py +++ b/python/sglang/srt/arg_groups/serving_hook.py @@ -851,11 +851,6 @@ def handle_multimodal_feature_transport(server_args: Any): raise ValueError( "--mm-feature-transport=cuda_vmm does not support pipeline parallelism." ) - if envs.SGLANG_RUST_SERVER.get(): - raise ValueError( - "--mm-feature-transport=cuda_vmm is not supported with " - "SGLANG_RUST_SERVER." - ) pool_budget_mb = envs.SGLANG_MM_FEATURE_CACHE_MB.get() handle_kind = "CUDA FABRIC" if cfg.nnodes > 1 else "POSIX FD" logger.info( diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 41ba165d9..891f831c5 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -686,15 +686,19 @@ class TestMultimodalFeatureTransport(CustomTestCase): handle_multimodal_feature_transport(server_args) @override_platform(is_cuda=True) - def test_cuda_vmm_rejects_rust_server(self): + def test_cuda_vmm_allows_rust_server(self): server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm") with ( + patch.dict(os.environ, {}, clear=False), envs.SGLANG_RUST_SERVER.override(True), - self.assertRaisesRegex(ValueError, "SGLANG_RUST_SERVER"), ): handle_multimodal_feature_transport(server_args) + self.assertEqual( + resolution_result(server_args, "mm_feature_transport"), "cuda_vmm" + ) + @override_platform(is_cuda=True) def test_cuda_vmm_rejects_pipeline_parallelism(self): server_args = ServerArgs(