From 9f56b471aaf528b3e98f0162e06513f0b34094e6 Mon Sep 17 00:00:00 2001 From: psaab Date: Mon, 16 Mar 2026 19:59:49 -0700 Subject: [PATCH] [Network] Use `NetworkAddress` for `dist_init_method` and loopback fallbacks (#20657) Co-authored-by: hnyls2002 --- .../runtime/managers/gpu_worker.py | 5 ++- .../disaggregation/ascend/transfer_engine.py | 5 ++- .../sglang/srt/disaggregation/common/conn.py | 2 +- .../srt/disaggregation/encode_grpc_server.py | 9 +++-- .../srt/disaggregation/encode_server.py | 23 +++++++----- .../srt/managers/data_parallel_controller.py | 3 +- .../sglang/srt/model_executor/model_runner.py | 30 +++++++++------- python/sglang/srt/utils/network.py | 36 +++++++++++++------ test/registered/utils/test_network_address.py | 10 +++--- 9 files changed, 79 insertions(+), 44 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index eb225a43c..466fa74b6 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -57,6 +57,7 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import ( PerformanceLogger, capture_memory_snapshot, ) +from sglang.srt.utils.network import NetworkAddress logger = init_logger(__name__) @@ -106,7 +107,9 @@ class GPUWorker: ring_degree=self.server_args.ring_degree, sp_size=self.server_args.sp_degree, dp_size=self.server_args.dp_size, - distributed_init_method=f"tcp://127.0.0.1:{self.master_port}", + distributed_init_method=NetworkAddress( + "127.0.0.1", self.master_port + ).to_tcp(), dist_timeout=self.server_args.dist_timeout, ) diff --git a/python/sglang/srt/disaggregation/ascend/transfer_engine.py b/python/sglang/srt/disaggregation/ascend/transfer_engine.py index 53c402df1..5cef34ae9 100644 --- a/python/sglang/srt/disaggregation/ascend/transfer_engine.py +++ b/python/sglang/srt/disaggregation/ascend/transfer_engine.py @@ -8,6 +8,7 @@ from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( MooncakeTransferEngine, ) +from sglang.srt.utils.network import NetworkAddress try: from memfabric_hybrid import TransferEngine @@ -47,7 +48,9 @@ class AscendTransferEngine(MooncakeTransferEngine): else: logger.error(f"Unsupported DisaggregationMode: {disaggregation_mode}") raise ValueError(f"Unsupported DisaggregationMode: {disaggregation_mode}") - self.session_id = f"{self.hostname}:{self.engine.get_rpc_port()}" + self.session_id = NetworkAddress( + self.hostname, self.engine.get_rpc_port() + ).to_host_port_str() self.initialize() def initialize(self) -> None: diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 1e38ba747..7b93c26c1 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -238,7 +238,7 @@ class CommonKVManager(BaseKVManager): """Register prefill server info to bootstrap server via HTTP POST.""" if self.dist_init_addr: # Multi-node case: bootstrap server's host is dist_init_addr - host = NetworkAddress.parse(self.dist_init_addr).host + host = NetworkAddress.parse(self.dist_init_addr).resolved().host else: # Single-node case: bootstrap server's host is the same as http server's host host = self.bootstrap_host diff --git a/python/sglang/srt/disaggregation/encode_grpc_server.py b/python/sglang/srt/disaggregation/encode_grpc_server.py index f52cd4e86..b7d2e5f03 100644 --- a/python/sglang/srt/disaggregation/encode_grpc_server.py +++ b/python/sglang/srt/disaggregation/encode_grpc_server.py @@ -29,7 +29,7 @@ from sglang.srt.disaggregation.encode_server import ( from sglang.srt.managers.schedule_batch import Modality from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import random_uuid -from sglang.srt.utils.network import get_zmq_socket +from sglang.srt.utils.network import NetworkAddress, get_zmq_socket logger = logging.getLogger(__name__) SGLangEncoderServicer = sglang_encoder_pb2_grpc.SglangEncoderServicer @@ -212,9 +212,12 @@ async def serve_grpc_encoder(server_args: ServerArgs): port_args = PortArgs.init_new(server_args) if server_args.dist_init_addr: - dist_init_method = f"tcp://{server_args.dist_init_addr}" + na = NetworkAddress.parse(server_args.dist_init_addr) + dist_init_method = na.to_tcp() else: - dist_init_method = f"tcp://127.0.0.1:{port_args.nccl_port}" + dist_init_method = NetworkAddress( + server_args.host or "127.0.0.1", port_args.nccl_port + ).to_tcp() send_sockets: List[zmq.Socket] = [] for rank in range(1, server_args.tp_size): diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 12801b8b2..4da24e522 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -49,7 +49,12 @@ from sglang.srt.utils import ( load_video, random_uuid, ) -from sglang.srt.utils.network import config_socket, get_local_ip_auto, get_zmq_socket +from sglang.srt.utils.network import ( + NetworkAddress, + config_socket, + get_local_ip_auto, + get_zmq_socket, +) logger = logging.getLogger(__name__) @@ -1002,11 +1007,10 @@ class MMEncoder: mm_data.embedding = None # Send ack/data - endpoint = ( - f"tcp://{url}" - if url is not None - else f"tcp://{prefill_host}:{embedding_port}" - ) + if url is not None: + endpoint = NetworkAddress.parse(url).to_tcp() + else: + endpoint = NetworkAddress(prefill_host, embedding_port).to_tcp() logger.info(f"{endpoint = }") # Serialize data @@ -1303,9 +1307,12 @@ def launch_server(server_args: ServerArgs): ipc_path_prefix = random_uuid() port_args = PortArgs.init_new(server_args) if server_args.dist_init_addr: - dist_init_method = f"tcp://{server_args.dist_init_addr}" + na = NetworkAddress.parse(server_args.dist_init_addr) + dist_init_method = na.to_tcp() else: - dist_init_method = f"tcp://127.0.0.1:{port_args.nccl_port}" + dist_init_method = NetworkAddress( + server_args.host or "127.0.0.1", port_args.nccl_port + ).to_tcp() for rank in range(1, server_args.tp_size): schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" send_sockets.append( diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 7aeaafdbb..9b3f4f270 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -288,7 +288,8 @@ class DataParallelController: # Determine the endpoint for inter-node communication if server_args.dist_init_addr is None: na = NetworkAddress( - "127.0.0.1", server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA + server_args.host or "127.0.0.1", + server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA, ) else: na = NetworkAddress.parse(server_args.dist_init_addr) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index d85a7d1bf..e73ea7bac 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -176,7 +176,7 @@ from sglang.srt.utils import ( set_cuda_arch, slow_rank_detector, ) -from sglang.srt.utils.network import get_local_ip_auto +from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.offloader import ( create_offloader_from_server_args, @@ -672,9 +672,9 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.remote_instance_transfer_engine.initialize( local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.get() ) - self.remote_instance_transfer_engine_session_id = ( - f"{local_ip}:{self.remote_instance_transfer_engine.get_rpc_port()}" - ) + self.remote_instance_transfer_engine_session_id = NetworkAddress( + local_ip, self.remote_instance_transfer_engine.get_rpc_port() + ).to_host_port_str() def model_specific_adjustment(self): server_args = self.server_args @@ -774,9 +774,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): if dist_init_method_override: dist_init_method = dist_init_method_override elif self.server_args.dist_init_addr: - dist_init_method = f"tcp://{self.server_args.dist_init_addr}" + na = NetworkAddress.parse(self.server_args.dist_init_addr) + dist_init_method = na.to_tcp() else: - dist_init_method = f"tcp://127.0.0.1:{self.dist_port}" + dist_init_method = NetworkAddress( + self.server_args.host or "127.0.0.1", self.dist_port + ).to_tcp() set_custom_all_reduce(not self.server_args.disable_custom_all_reduce) set_mscclpp_all_reduce(self.server_args.enable_mscclpp) set_torch_symm_mem_all_reduce(self.server_args.enable_torch_symm_mem) @@ -956,7 +959,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): == RemoteInstanceWeightLoaderBackend.NCCL ): if self.tp_rank == 0: - instance_ip = socket.gethostbyname(socket.gethostname()) + instance_ip = NetworkAddress.resolve_host(socket.gethostname()) t = threading.Thread( target=trigger_init_weights_send_group_for_remote_instance_request, args=( @@ -1234,9 +1237,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): success = False message = "" try: + na = NetworkAddress(master_address, group_port) self._weights_send_group[group_name] = init_custom_process_group( backend=backend, - init_method=f"tcp://{master_address}:{group_port}", + init_method=na.to_tcp(), world_size=world_size, rank=group_rank, group_name=group_name, @@ -1244,9 +1248,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) dist.barrier(group=self._weights_send_group[group_name]) success = True - message = ( - f"Succeeded to init group through {master_address}:{group_port} group." - ) + message = f"Succeeded to init group through {na.to_host_port_str()} group." except Exception as e: message = f"Failed to init group: {e}." logger.error(message) @@ -1281,6 +1283,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): torch.cuda.empty_cache() success = False + na = NetworkAddress(master_address, group_port) message = "" try: for _, weights in self.model.named_parameters(): @@ -1290,7 +1293,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): group=send_group, ) success = True - message = f"Succeeded to send weights through {master_address}:{group_port} {group_name}." + message = f"Succeeded to send weights through {na.to_host_port_str()} {group_name}." except Exception as e: message = f"Failed to send weights: {e}." logger.error(message) @@ -1333,9 +1336,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) try: + na = NetworkAddress(master_address, master_port) self._model_update_group[group_name] = init_custom_process_group( backend=backend, - init_method=f"tcp://{master_address}:{master_port}", + init_method=na.to_tcp(), world_size=world_size, rank=rank, group_name=group_name, diff --git a/python/sglang/srt/utils/network.py b/python/sglang/srt/utils/network.py index 9d982c4db..c374c9535 100644 --- a/python/sglang/srt/utils/network.py +++ b/python/sglang/srt/utils/network.py @@ -416,6 +416,11 @@ class NetworkAddress: host: str port: int + def __post_init__(self): + # Auto-strip IPv6 brackets so callers can pass "[::1]" or "::1" + if self.host.startswith("[") and self.host.endswith("]"): + object.__setattr__(self, "host", self.host[1:-1]) + @property def is_ipv6(self) -> bool: return _is_ipv6(self.host) @@ -436,6 +441,26 @@ class NetworkAddress: """``host:port`` string for gRPC listen address, session IDs, logs.""" return f"{_wrap(self.host)}:{self.port}" + @staticmethod + def resolve_host(host: str) -> str: + """Return *host* as-is if it's an IP, otherwise DNS-resolve to one.""" + try: + ipaddress.ip_address(host) + return host + except ValueError: + pass + try: + return socket.getaddrinfo( + host, None, socket.AF_UNSPEC, 0, 0, socket.AI_ADDRCONFIG + )[0][4][0] + except socket.gaierror as e: + raise ValueError(f"Cannot resolve host {host!r}: {e}") from e + + def resolved(self) -> NetworkAddress: + """DNS-resolve hostname to IP; return self if already an IP.""" + ip = self.resolve_host(self.host) + return self if ip == self.host else NetworkAddress(ip, self.port) + def to_bind_tuple(self) -> Tuple[str, int]: """Raw ``(host, port)`` tuple for ``socket.bind()`` / ``socket.connect()``. @@ -492,17 +517,6 @@ class NetworkAddress: ) return NetworkAddress(host, _parse_port(port_str)) - @staticmethod - def from_parts(host: str, port: int) -> NetworkAddress: - """Create from separate host and port, stripping brackets if present. - - Useful when the host may come from user input that already has - brackets (e.g. ``[::1]``). - """ - if host.startswith("[") and host.endswith("]"): - host = host[1:-1] - return NetworkAddress(host, port) - def __str__(self) -> str: return self.to_host_port_str() diff --git a/test/registered/utils/test_network_address.py b/test/registered/utils/test_network_address.py index 7230c83f6..922866f26 100644 --- a/test/registered/utils/test_network_address.py +++ b/test/registered/utils/test_network_address.py @@ -179,23 +179,23 @@ class TestNetworkAddressParseErrors(unittest.TestCase): NetworkAddress.parse(":8000") -class TestNetworkAddressFromParts(unittest.TestCase): +class TestNetworkAddressBracketStripping(unittest.TestCase): def test_strip_brackets(self): - na = NetworkAddress.from_parts("[::1]", 8000) + na = NetworkAddress("[::1]", 8000) self.assertEqual(na.host, "::1") self.assertTrue(na.is_ipv6) def test_no_brackets(self): - na = NetworkAddress.from_parts("::1", 8000) + na = NetworkAddress("::1", 8000) self.assertEqual(na.host, "::1") def test_ipv4_passthrough(self): - na = NetworkAddress.from_parts("127.0.0.1", 30000) + na = NetworkAddress("127.0.0.1", 30000) self.assertEqual(na.host, "127.0.0.1") self.assertFalse(na.is_ipv6) def test_hostname_passthrough(self): - na = NetworkAddress.from_parts("myhost", 30000) + na = NetworkAddress("myhost", 30000) self.assertEqual(na.host, "myhost")