From 597d17dd18af037ec697a784494a6e5bfff58fa2 Mon Sep 17 00:00:00 2001 From: Chanh Nguyen Date: Sun, 15 Feb 2026 08:32:47 -0800 Subject: [PATCH] Use ephemeral nccl port via get_free_port() (#18009) Co-authored-by: Chanh Nguyen --- python/sglang/srt/server_args.py | 27 +++---- test/registered/core/test_server_args.py | 92 +++++++++--------------- 2 files changed, 45 insertions(+), 74 deletions(-) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8c187d15f..2c9f79268 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -43,6 +43,7 @@ from sglang.srt.utils.common import ( get_device_memory_capacity, get_device_name, get_device_sm, + get_free_port, get_int_env_var, get_quantization_config, is_blackwell_supported, @@ -52,7 +53,6 @@ from sglang.srt.utils.common import ( is_hopper_with_cuda_12_3, is_no_spec_infer_or_topk_one, is_npu, - is_port_available, is_remote_url, is_sm90_supported, is_sm100_supported, @@ -1135,9 +1135,9 @@ class ServerArgs: # suggest them to be explicit about kv_cache_dtype to avoid surprises if (user_set_prefill or user_set_decode) and self.kv_cache_dtype == "auto": logger.warning( - f"When specifying --nsa-prefill-backend or --nsa-decode-backend, " - f"you should also explicitly set --kv-cache-dtype (e.g., 'fp8_e4m3' or 'bfloat16'). " - f"DeepSeek V3.2 defaults to FP8 KV cache which may not be compatible with all backends." + "When specifying --nsa-prefill-backend or --nsa-decode-backend, " + "you should also explicitly set --kv-cache-dtype (e.g., 'fp8_e4m3' or 'bfloat16'). " + "DeepSeek V3.2 defaults to FP8 KV cache which may not be compatible with all backends." ) if self.kv_cache_dtype == "auto": @@ -1216,7 +1216,7 @@ class ServerArgs: if not is_npu(): # CUDA or ROCm GPU if self.enable_nsa_prefill_context_parallel: logger.warning( - f"Context parallel feature is still under experiment. It has only been verified on Hopper platform." + "Context parallel feature is still under experiment. It has only been verified on Hopper platform." ) if self.nsa_prefill_cp_mode == "in-seq-split": # TODO Supports moe_dense_tp_size != 1, kv cache dtype = "fp8",moe_a2a_backend non-deepep and cross-machine operation . @@ -1226,7 +1226,7 @@ class ServerArgs: self.ep_size = self.tp_size self.kv_cache_dtype = "bf16" logger.warning( - f"For in-seq split mode, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, kv_cache_dtype == bf16, batch_size == 1" + "For in-seq split mode, we have the following restrictions: moe_dense_tp_size == 1, moe_a2a_backend == deepep, ep_size == tp_size, kv_cache_dtype == bf16, batch_size == 1" ) else: self.enable_dp_attention = True @@ -1663,7 +1663,7 @@ class ServerArgs: if envs.SGLANG_EMBEDDINGS_SPARSE_HEAD.is_set(): self.disable_overlap_schedule = True logger.warning( - f"Overlap scheduler is disabled when using sparse head for embedding model." + "Overlap scheduler is disabled when using sparse head for embedding model." ) # TRTLLM AllReduce Fusion supports SM90/100, enable it by default @@ -5500,7 +5500,7 @@ class ServerArgs: def validate_transfer_engine(self): if importlib.util.find_spec("mooncake.engine") is None: logger.warning( - f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend." + "Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend." ) return False elif self.enable_memory_saver: @@ -5603,14 +5603,7 @@ class PortArgs: worker_ports: Optional[List[int]] = None, ) -> PortArgs: if server_args.nccl_port is None: - nccl_port = server_args.port + random.randint(100, 1000) - while True: - if is_port_available(nccl_port): - break - if nccl_port < 60000: - nccl_port += 42 - else: - nccl_port -= 43 + nccl_port = get_free_port() else: nccl_port = server_args.nccl_port @@ -5671,7 +5664,7 @@ class PortArgs: # Skip check when using worker_ports since the port is already bound by our ZMQ socket if dp_rank is None or worker_ports is None: wait_port_available(scheduler_input_port, "scheduler_input_port") - except ValueError as e: + except ValueError: logger.exception( f"Port is already in use. {dist_init_port=} {port_base=} {detokenizer_port=} {nccl_port=} {scheduler_input_port=}" ) diff --git a/test/registered/core/test_server_args.py b/test/registered/core/test_server_args.py index 91e90a3e4..d7382d0dd 100644 --- a/test/registered/core/test_server_args.py +++ b/test/registered/core/test_server_args.py @@ -1,6 +1,6 @@ import json import unittest -from unittest.mock import patch +from unittest.mock import MagicMock, patch from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci @@ -44,10 +44,29 @@ class TestLoadBalanceMethod(unittest.TestCase): class TestPortArgs(unittest.TestCase): - @patch("sglang.srt.server_args.is_port_available") + @patch("sglang.srt.server_args.get_free_port") @patch("sglang.srt.server_args.tempfile.NamedTemporaryFile") - def test_init_new_standard_case(self, mock_temp_file, mock_is_port_available): - mock_is_port_available.return_value = True + def test_init_new_with_nccl_port_none(self, mock_temp_file, mock_get_free_port): + """Test that get_free_port() is called when nccl_port is None""" + mock_temp_file.return_value.name = "temp_file" + mock_get_free_port.return_value = 45678 # Mock ephemeral port + + # Use MagicMock here to verify get_free_port is called + server_args = MagicMock() + server_args.nccl_port = None + server_args.enable_dp_attention = False + server_args.tokenizer_worker_num = 1 + + port_args = PortArgs.init_new(server_args) + + # Verify get_free_port was called + mock_get_free_port.assert_called_once() + + # Verify the returned port is used + self.assertEqual(port_args.nccl_port, 45678) + + @patch("sglang.srt.server_args.tempfile.NamedTemporaryFile") + def test_init_new_standard_case(self, mock_temp_file): mock_temp_file.return_value.name = "temp_file" server_args = ServerArgs(model_path="dummy") @@ -62,9 +81,7 @@ class TestPortArgs(unittest.TestCase): self.assertTrue(port_args.detokenizer_ipc_name.startswith("ipc://")) self.assertIsInstance(port_args.nccl_port, int) - @patch("sglang.srt.server_args.is_port_available") - def test_init_new_with_single_node_dp_attention(self, mock_is_port_available): - mock_is_port_available.return_value = True + def test_init_new_with_single_node_dp_attention(self): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 @@ -82,10 +99,7 @@ class TestPortArgs(unittest.TestCase): self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://127.0.0.1:")) self.assertIsInstance(port_args.nccl_port, int) - @patch("sglang.srt.server_args.is_port_available") - def test_init_new_with_dp_rank(self, mock_is_port_available): - mock_is_port_available.return_value = True - + def test_init_new_with_dp_rank(self): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -102,13 +116,9 @@ class TestPortArgs(unittest.TestCase): self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://192.168.1.1:")) self.assertIsInstance(port_args.nccl_port, int) - @patch("sglang.srt.server_args.is_port_available") - def test_init_new_with_ipv4_address(self, mock_is_port_available): - mock_is_port_available.return_value = True - + def test_init_new_with_ipv4_address(self): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 - server_args.nccl_port = None server_args.enable_dp_attention = True @@ -124,10 +134,7 @@ class TestPortArgs(unittest.TestCase): self.assertTrue(port_args.detokenizer_ipc_name.startswith("tcp://192.168.1.1:")) self.assertIsInstance(port_args.nccl_port, int) - @patch("sglang.srt.server_args.is_port_available") - def test_init_new_with_malformed_ipv4_address(self, mock_is_port_available): - mock_is_port_available.return_value = True - + def test_init_new_with_malformed_ipv4_address(self): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -143,12 +150,7 @@ class TestPortArgs(unittest.TestCase): "please provide --dist-init-addr as host:port", str(context.exception) ) - @patch("sglang.srt.server_args.is_port_available") - def test_init_new_with_malformed_ipv4_address_invalid_port( - self, mock_is_port_available - ): - mock_is_port_available.return_value = True - + def test_init_new_with_malformed_ipv4_address_invalid_port(self): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -157,16 +159,11 @@ class TestPortArgs(unittest.TestCase): server_args.nnodes = 2 server_args.dist_init_addr = "192.168.1.1:abc" - with self.assertRaises(ValueError) as context: + with self.assertRaises(ValueError): PortArgs.init_new(server_args) - @patch("sglang.srt.server_args.is_port_available") @patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True) - def test_init_new_with_ipv6_address( - self, mock_is_valid_ipv6, mock_is_port_available - ): - mock_is_port_available.return_value = True - + def test_init_new_with_ipv6_address(self, mock_is_valid_ipv6): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -186,13 +183,8 @@ class TestPortArgs(unittest.TestCase): ) self.assertIsInstance(port_args.nccl_port, int) - @patch("sglang.srt.server_args.is_port_available") @patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=False) - def test_init_new_with_invalid_ipv6_address( - self, mock_is_valid_ipv6, mock_is_port_available - ): - mock_is_port_available.return_value = True - + def test_init_new_with_invalid_ipv6_address(self, mock_is_valid_ipv6): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -206,12 +198,7 @@ class TestPortArgs(unittest.TestCase): self.assertIn("invalid IPv6 address", str(context.exception)) - @patch("sglang.srt.server_args.is_port_available") - def test_init_new_with_malformed_ipv6_address_missing_bracket( - self, mock_is_port_available - ): - mock_is_port_available.return_value = True - + def test_init_new_with_malformed_ipv6_address_missing_bracket(self): server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -225,13 +212,10 @@ class TestPortArgs(unittest.TestCase): self.assertIn("invalid IPv6 address format", str(context.exception)) - @patch("sglang.srt.server_args.is_port_available") @patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True) def test_init_new_with_malformed_ipv6_address_missing_port( - self, mock_is_valid_ipv6, mock_is_port_available + self, mock_is_valid_ipv6 ): - mock_is_port_available.return_value = True - server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -247,13 +231,10 @@ class TestPortArgs(unittest.TestCase): "a port must be specified in IPv6 address", str(context.exception) ) - @patch("sglang.srt.server_args.is_port_available") @patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True) def test_init_new_with_malformed_ipv6_address_invalid_port( - self, mock_is_valid_ipv6, mock_is_port_available + self, mock_is_valid_ipv6 ): - mock_is_port_available.return_value = True - server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None @@ -267,13 +248,10 @@ class TestPortArgs(unittest.TestCase): self.assertIn("invalid port in IPv6 address", str(context.exception)) - @patch("sglang.srt.server_args.is_port_available") @patch("sglang.srt.server_args.is_valid_ipv6_address", return_value=True) def test_init_new_with_malformed_ipv6_address_wrong_separator( - self, mock_is_valid_ipv6, mock_is_port_available + self, mock_is_valid_ipv6 ): - mock_is_port_available.return_value = True - server_args = ServerArgs(model_path="dummy") server_args.port = 30000 server_args.nccl_port = None