[Utils] Move network/socket utilities from common.py to network.py (#20646)

This commit is contained in:
Liangsheng Yin
2026-03-15 20:35:24 -07:00
committed by GitHub
parent afc71bae3a
commit f0458e0b49
28 changed files with 418 additions and 422 deletions
@@ -382,7 +382,7 @@ class Runtime:
# client code without installing SRT server and its dependency if they want. # client code without installing SRT server and its dependency if they want.
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_port_available from sglang.srt.utils.network import is_port_available
self.server_args = ServerArgs(*args, log_level=log_level, **kwargs) self.server_args = ServerArgs(*args, log_level=log_level, **kwargs)
@@ -13,7 +13,7 @@ from sglang.srt.disaggregation.mooncake.conn import (
MooncakeKVReceiver, MooncakeKVReceiver,
MooncakeKVSender, MooncakeKVSender,
) )
from sglang.srt.utils import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -35,11 +35,11 @@ from sglang.srt.layers.dp_attention import (
get_attention_tp_size, get_attention_tp_size,
) )
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils.network import (
NetworkAddress,
get_local_ip_auto, get_local_ip_auto,
get_zmq_socket_on_host, get_zmq_socket_on_host,
) )
from sglang.srt.utils.network import NetworkAddress
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -27,7 +27,8 @@ from sglang.srt.disaggregation.encode_server import (
launch_encoder, launch_encoder,
) )
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import get_zmq_socket, random_uuid from sglang.srt.utils import random_uuid
from sglang.srt.utils.network import get_zmq_socket
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
SGLangEncoderServicer = sglang_encoder_pb2_grpc.SglangEncoderServicer SGLangEncoderServicer = sglang_encoder_pb2_grpc.SglangEncoderServicer
@@ -25,8 +25,9 @@ from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReq
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData, get_local_ip_auto, get_zmq_socket_on_host from sglang.srt.utils import ImageData
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
from sglang.srt.utils.network import get_local_ip_auto, get_zmq_socket_on_host
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -43,14 +43,12 @@ from sglang.srt.server_args import (
set_global_server_args_for_scheduler, set_global_server_args_for_scheduler,
) )
from sglang.srt.utils import ( from sglang.srt.utils import (
config_socket,
get_local_ip_auto,
get_zmq_socket,
load_audio, load_audio,
load_image, load_image,
load_video, load_video,
random_uuid, random_uuid,
) )
from sglang.srt.utils.network import config_socket, get_local_ip_auto, get_zmq_socket
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -37,11 +37,8 @@ from sglang.srt.disaggregation.utils import (
filter_kv_indices_for_cp_rank, filter_kv_indices_for_cp_rank,
) )
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import ( from sglang.srt.utils.common import get_int_env_var
get_int_env_var, from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
get_local_ip_auto,
)
from sglang.srt.utils.network import NetworkAddress
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
MORI_GUARD = b"MoriMsgGuard" MORI_GUARD = b"MoriMsgGuard"
@@ -4,8 +4,7 @@ import os
from typing import List, Optional from typing import List, Optional
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.utils import get_free_port from sglang.srt.utils.network import NetworkAddress, get_free_port
from sglang.srt.utils.network import NetworkAddress
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -16,11 +16,7 @@ from torch.distributed import ProcessGroup
from zmq import IPV6 # type: ignore from zmq import IPV6 # type: ignore
from zmq import SUB, SUBSCRIBE, XPUB, XPUB_VERBOSE, Context # type: ignore from zmq import SUB, SUBSCRIBE, XPUB, XPUB_VERBOSE, Context # type: ignore
from sglang.srt.utils import ( from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto, get_open_port
get_local_ip_auto,
get_open_port,
)
from sglang.srt.utils.network import NetworkAddress
# SGLANG_RINGBUFFER_WARNING_INTERVAL can be set to 60 # SGLANG_RINGBUFFER_WARNING_INTERVAL can be set to 60
SGLANG_RINGBUFFER_WARNING_INTERVAL = int( SGLANG_RINGBUFFER_WARNING_INTERVAL = int(
@@ -48,7 +48,6 @@ from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
get_current_device_stream_fast, get_current_device_stream_fast,
get_int_env_var, get_int_env_var,
get_local_ip_auto,
is_cpu, is_cpu,
is_cuda_alike, is_cuda_alike,
is_hip, is_hip,
@@ -58,6 +57,7 @@ from sglang.srt.utils import (
is_xpu, is_xpu,
) )
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
from sglang.srt.utils.network import get_local_ip_auto
_is_npu = is_npu() _is_npu = is_npu()
_is_cpu = is_cpu() _is_cpu = is_cpu()
@@ -14,7 +14,7 @@ from sglang.srt.environ import envs
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.managers.io_struct import UpdateExpertBackupReq from sglang.srt.managers.io_struct import UpdateExpertBackupReq
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
PORT_BASE = envs.SGLANG_BACKUP_PORT_BASE.get() PORT_BASE = envs.SGLANG_BACKUP_PORT_BASE.get()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -17,7 +17,7 @@ from sglang.srt.server_args import (
ServerArgs, ServerArgs,
set_global_server_args_for_scheduler, set_global_server_args_for_scheduler,
) )
from sglang.srt.utils import get_local_ip_auto from sglang.srt.utils.network import get_local_ip_auto
PORT_BASE = envs.SGLANG_BACKUP_PORT_BASE.get() PORT_BASE = envs.SGLANG_BACKUP_PORT_BASE.get()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
+1 -1
View File
@@ -88,7 +88,6 @@ from sglang.srt.utils import (
assert_pkg_version, assert_pkg_version,
configure_logger, configure_logger,
get_bool_env_var, get_bool_env_var,
get_zmq_socket,
is_cuda, is_cuda,
kill_process_tree, kill_process_tree,
launch_dummy_health_check_server, launch_dummy_health_check_server,
@@ -97,6 +96,7 @@ from sglang.srt.utils import (
set_prometheus_multiproc_dir, set_prometheus_multiproc_dir,
set_ulimit, set_ulimit,
) )
from sglang.srt.utils.network import get_zmq_socket
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.version import __version__ from sglang.version import __version__
@@ -47,13 +47,11 @@ from sglang.srt.server_args import (
) )
from sglang.srt.utils import numa_utils from sglang.srt.utils import numa_utils
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
bind_port,
configure_logger, configure_logger,
get_zmq_socket,
kill_itself_when_parent_died, kill_itself_when_parent_died,
maybe_reindex_device_id, maybe_reindex_device_id,
) )
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress, bind_port, get_zmq_socket
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.srt.utils.watchdog import Watchdog from sglang.srt.utils.watchdog import Watchdog
from sglang.utils import TypeBasedDispatcher, get_exception_traceback from sglang.utils import TypeBasedDispatcher, get_exception_traceback
@@ -39,10 +39,10 @@ from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils import (
configure_logger, configure_logger,
freeze_gc, freeze_gc,
get_zmq_socket,
kill_itself_when_parent_died, kill_itself_when_parent_died,
) )
from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.srt.utils.hf_transformers_utils import get_tokenizer
from sglang.srt.utils.network import get_zmq_socket
from sglang.srt.utils.watchdog import Watchdog from sglang.srt.utils.watchdog import Watchdog
from sglang.utils import ( from sglang.utils import (
TypeBasedDispatcher, TypeBasedDispatcher,
@@ -47,7 +47,8 @@ from sglang.srt.managers.io_struct import (
from sglang.srt.managers.tokenizer_communicator_mixin import _Communicator from sglang.srt.managers.tokenizer_communicator_mixin import _Communicator
from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import get_zmq_socket, kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.network import get_zmq_socket
from sglang.utils import get_exception_traceback from sglang.utils import get_exception_traceback
if TYPE_CHECKING: if TYPE_CHECKING:
+1 -1
View File
@@ -201,7 +201,6 @@ from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
get_int_env_var, get_int_env_var,
get_numa_node, get_numa_node,
get_zmq_socket,
is_mps, is_mps,
kill_itself_when_parent_died, kill_itself_when_parent_died,
numa_bind_to_node, numa_bind_to_node,
@@ -217,6 +216,7 @@ from sglang.srt.utils.hf_transformers_utils import (
get_tokenizer, get_tokenizer,
get_tokenizer_from_processor, get_tokenizer_from_processor,
) )
from sglang.srt.utils.network import get_zmq_socket
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.utils import TypeBasedDispatcher, get_exception_traceback from sglang.utils import TypeBasedDispatcher, get_exception_traceback
@@ -104,7 +104,6 @@ from sglang.srt.utils import (
freeze_gc, freeze_gc,
get_bool_env_var, get_bool_env_var,
get_or_create_event_loop, get_or_create_event_loop,
get_zmq_socket,
kill_process_tree, kill_process_tree,
) )
from sglang.srt.utils.aio_rwlock import RWLock from sglang.srt.utils.aio_rwlock import RWLock
@@ -113,6 +112,7 @@ from sglang.srt.utils.hf_transformers_utils import (
get_tokenizer, get_tokenizer,
get_tokenizer_from_processor, get_tokenizer_from_processor,
) )
from sglang.srt.utils.network import get_zmq_socket
from sglang.srt.utils.request_logger import RequestLogger from sglang.srt.utils.request_logger import RequestLogger
from sglang.srt.utils.watchdog import Watchdog from sglang.srt.utils.watchdog import Watchdog
from sglang.utils import TypeBasedDispatcher, get_exception_traceback from sglang.utils import TypeBasedDispatcher, get_exception_traceback
@@ -163,7 +163,6 @@ from sglang.srt.utils import (
enable_show_time_cost, enable_show_time_cost,
get_available_gpu_memory, get_available_gpu_memory,
get_cpu_ids_by_node, get_cpu_ids_by_node,
get_local_ip_auto,
init_custom_process_group, init_custom_process_group,
is_hip, is_hip,
is_host_cpu_arm64, is_host_cpu_arm64,
@@ -177,6 +176,7 @@ from sglang.srt.utils import (
set_cuda_arch, set_cuda_arch,
slow_rank_detector, slow_rank_detector,
) )
from sglang.srt.utils.network import get_local_ip_auto
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
from sglang.srt.utils.offloader import ( from sglang.srt.utils.offloader import (
create_offloader_from_server_args, create_offloader_from_server_args,
+1 -3
View File
@@ -41,7 +41,6 @@ from sglang.srt.utils.common import (
get_device_memory_capacity, get_device_memory_capacity,
get_device_name, get_device_name,
get_device_sm, get_device_sm,
get_free_port,
get_int_env_var, get_int_env_var,
get_quantization_config, get_quantization_config,
is_blackwell_supported, is_blackwell_supported,
@@ -63,11 +62,10 @@ from sglang.srt.utils.common import (
nullable_str, nullable_str,
parse_connector_type, parse_connector_type,
torch_release, torch_release,
wait_port_available,
xpu_has_xmx_support, xpu_has_xmx_support,
) )
from sglang.srt.utils.hf_transformers_utils import check_gguf_file from sglang.srt.utils.hf_transformers_utils import check_gguf_file
from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.network import NetworkAddress, get_free_port, wait_port_available
from sglang.utils import is_in_ci from sglang.utils import is_in_ci
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
-376
View File
@@ -23,7 +23,6 @@ import functools
import importlib import importlib
import inspect import inspect
import io import io
import ipaddress
import itertools import itertools
import json import json
import logging import logging
@@ -36,7 +35,6 @@ import re
import resource import resource
import shutil import shutil
import signal import signal
import socket
import subprocess import subprocess
import sys import sys
import tempfile import tempfile
@@ -81,7 +79,6 @@ import requests
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import triton import triton
import zmq
from packaging import version as pkg_version from packaging import version as pkg_version
from PIL import Image from PIL import Image
from starlette.routing import Mount from starlette.routing import Mount
@@ -702,143 +699,6 @@ def set_random_seed(seed: int) -> None:
torch.cuda.manual_seed_all(seed) torch.cuda.manual_seed_all(seed)
def find_process_using_port(port: int) -> Optional[psutil.Process]:
for conn in psutil.net_connections(kind="inet"):
if conn.laddr.port == port:
try:
return psutil.Process(conn.pid)
except psutil.NoSuchProcess:
# It could happen by race condition (the proc dies when psutil.Process is called).
pass
return None
def wait_port_available(
port: int, port_name: str, timeout_s: int = 30, raise_exception: bool = True
) -> bool:
for i in range(timeout_s):
if is_port_available(port):
return True
if i > 10 and i % 5 == 0:
process = find_process_using_port(port)
if process is None:
logger.warning(
f"The port {port} is in use, but we could not find the process that uses it."
)
pid = process.pid
error_message = f"{port_name} is used by a process already. {process.name()=}' {process.cmdline()=} {process.status()=} {pid=}"
logger.info(
f"port {port} is in use. Waiting for {i} seconds for {port_name} to be available. {error_message}"
)
time.sleep(0.1)
if raise_exception:
raise ValueError(
f"{port_name} at {port} is not available in {timeout_s} seconds. {error_message}"
)
return False
def _get_addrinfos_for_bind(host=None, port=0):
"""Return deduplicated addrinfo tuples for binding (one per address family).
Args:
host: Bind address. None (with AI_PASSIVE) resolves to wildcard
addresses (0.0.0.0 / ::) suitable for accepting on all interfaces.
port: Port number. 0 lets the OS assign an available ephemeral port.
Flags:
AI_ADDRCONFIG — only return families actually configured on this host.
AI_PASSIVE — return wildcard addresses suitable for bind().
Falls back to AF_INET if getaddrinfo fails (e.g. DNS misconfiguration).
"""
try:
infos = socket.getaddrinfo(
host,
port,
socket.AF_UNSPEC,
socket.SOCK_STREAM,
0,
socket.AI_ADDRCONFIG | socket.AI_PASSIVE,
)
deduped = []
seen_families = set()
for info in infos:
if info[0] not in seen_families:
seen_families.add(info[0])
deduped.append(info)
# Prefer IPv4 so that callers without an explicit host get consistent
# behaviour across platforms (some OSes list IPv6 first).
deduped.sort(key=lambda x: (x[0] != socket.AF_INET,))
return deduped
except socket.gaierror:
fallback_host = "0.0.0.0" if host is None else host
return [(socket.AF_INET, socket.SOCK_STREAM, 0, "", (fallback_host, port))]
def try_bind_socket(host=None, port=0, *, reuse_addr=True, listen=False):
"""Bind a TCP socket on the first available address family (IPv4/IPv6).
Iterates over address families returned by _get_addrinfos_for_bind and
returns the first socket that successfully binds.
Args:
host: Bind address. None binds to all interfaces (0.0.0.0 / ::).
port: Port number. 0 lets the OS assign an available ephemeral port;
use sock.getsockname()[1] to retrieve the assigned port.
reuse_addr: Set SO_REUSEADDR to allow quick port reuse after close.
listen: Call listen(1) after bind, making the socket ready to accept.
Returns:
The bound socket. Caller is responsible for closing it.
Raises:
OSError: If bind fails on all configured address families.
"""
for family, socktype, proto, _, sockaddr in _get_addrinfos_for_bind(host, port):
sock = socket.socket(family, socktype, proto)
try:
if family == socket.AF_INET6:
sock.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
if reuse_addr:
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
sock.bind(sockaddr)
if listen:
sock.listen(1)
return sock
except (OSError, OverflowError):
sock.close()
raise OSError(f"Could not bind port {port} on any configured address family")
def is_port_available(port):
"""Return whether a port is available on all configured address families."""
try:
for family, socktype, proto, _, sockaddr in _get_addrinfos_for_bind(port=port):
sock = socket.socket(family, socktype, proto)
try:
if family == socket.AF_INET6:
sock.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
sock.bind(sockaddr)
finally:
sock.close()
return True
except (OSError, OverflowError):
return False
def get_free_port():
sock = try_bind_socket()
port = sock.getsockname()[1]
sock.close()
return port
def decode_video_base64(video_base64): def decode_video_base64(video_base64):
from PIL import Image from PIL import Image
@@ -1500,104 +1360,6 @@ def pytorch_profile(name, func, *args, data_size=-1):
return result return result
def get_zmq_socket(
context: zmq.Context,
socket_type: zmq.SocketType,
endpoint: Optional[str] = None,
bind: bool = True,
) -> Union[zmq.Socket, Tuple[int, zmq.Socket]]:
"""Create and configure a ZeroMQ socket.
Args:
context: ZeroMQ context to create the socket from.
socket_type: Type of ZeroMQ socket to create.
endpoint: Optional endpoint to bind/connect to. If None, binds to a random TCP port.
bind: Whether to bind (True) or connect (False) to the endpoint. Ignored if endpoint is None.
Returns:
If endpoint is None: Tuple of (port, socket) where port is the randomly assigned TCP port.
If endpoint is provided: The configured ZeroMQ socket.
"""
socket = context.socket(socket_type)
if endpoint is None:
# Bind to random TCP port
config_socket(socket, socket_type)
port = socket.bind_to_random_port("tcp://*")
return port, socket
else:
# Handle IPv6 if endpoint contains brackets
if endpoint.find("[") != -1:
socket.setsockopt(zmq.IPV6, 1)
config_socket(socket, socket_type)
if bind:
socket.bind(endpoint)
else:
socket.connect(endpoint)
return socket
def get_zmq_socket_on_host(
context: zmq.Context,
socket_type: zmq.SocketType,
host: Optional[str] = None,
) -> Tuple[int, zmq.Socket]:
"""Create and configure a ZeroMQ socket.
Args:
context: ZeroMQ context to create the socket from.
socket_type: Type of ZeroMQ socket to create.
host: Optional host to bind/connect to, without "tcp://" prefix. If None, binds to "tcp://*".
Returns:
Tuple of (port, socket) where port is the randomly assigned TCP port.
"""
socket = context.socket(socket_type)
# Bind to random TCP port, auto-wrapping IPv6 and setting zmq.IPV6 flag
config_socket(socket, socket_type)
if host:
if is_valid_ipv6_address(host):
socket.setsockopt(zmq.IPV6, 1)
bind_host = f"tcp://[{host}]"
else:
bind_host = f"tcp://{host}"
else:
bind_host = "tcp://*"
port = socket.bind_to_random_port(bind_host)
return port, socket
def config_socket(socket, socket_type: zmq.SocketType):
mem = psutil.virtual_memory()
total_mem = mem.total / 1024**3
available_mem = mem.available / 1024**3
if total_mem > 32 and available_mem > 16:
buf_size = int(0.5 * 1024**3)
else:
buf_size = -1
def set_send_opt():
socket.setsockopt(zmq.SNDHWM, 0)
socket.setsockopt(zmq.SNDBUF, buf_size)
def set_recv_opt():
socket.setsockopt(zmq.RCVHWM, 0)
socket.setsockopt(zmq.RCVBUF, buf_size)
if socket_type == zmq.PUSH:
set_send_opt()
elif socket_type == zmq.PULL:
set_recv_opt()
elif socket_type in [zmq.DEALER, zmq.REQ, zmq.REP]:
set_send_opt()
set_recv_opt()
else:
raise ValueError(f"Unsupported socket type: {socket_type}")
def dump_to_file(dirpath, name, value): def dump_to_file(dirpath, name, value):
from sglang.srt.distributed import get_tensor_model_parallel_rank from sglang.srt.distributed import get_tensor_model_parallel_rank
@@ -1766,11 +1528,6 @@ def _get_fastapi_request_path(request) -> Tuple[str, bool]:
return request.url.path, False return request.url.path, False
def bind_port(port):
"""Bind to a specific port, assuming it's available."""
return try_bind_socket(port=port, listen=True)
def get_amdgpu_memory_capacity(): def get_amdgpu_memory_capacity():
try: try:
# Run rocm-smi and capture the output # Run rocm-smi and capture the output
@@ -2707,29 +2464,6 @@ def _configure_uvicorn_access_log_filter(
filters_list.append(filter_name) filters_list.append(filter_name)
def get_open_port() -> int:
port = os.getenv("SGLANG_PORT")
if port is not None:
port = int(port)
while True:
if is_port_available(port):
return port
logger.info("Port %d is already in use, trying port %d", port, port + 1)
port += 1
sock = try_bind_socket()
port = sock.getsockname()[1]
sock.close()
return port
def is_valid_ipv6_address(address: str) -> bool:
try:
ipaddress.IPv6Address(address)
return True
except ValueError:
return False
def launch_dummy_health_check_server(host, port, enable_metrics): def launch_dummy_health_check_server(host, port, enable_metrics):
import asyncio import asyncio
@@ -2960,116 +2694,6 @@ def bind_or_assign(target, source):
return source return source
def get_local_ip_by_nic(interface: str = None) -> Optional[str]:
if not (interface := interface or os.environ.get("SGLANG_LOCAL_IP_NIC", None)):
return None
try:
import netifaces
except ImportError as e:
raise ImportError(
"Environment variable SGLANG_LOCAL_IP_NIC requires package netifaces, please install it through 'pip install netifaces'"
) from e
try:
addresses = netifaces.ifaddresses(interface)
if netifaces.AF_INET in addresses:
for addr_info in addresses[netifaces.AF_INET]:
ip = addr_info.get("addr")
if ip and ip != "127.0.0.1" and ip != "0.0.0.0":
return ip
if netifaces.AF_INET6 in addresses:
for addr_info in addresses[netifaces.AF_INET6]:
ip = addr_info.get("addr")
if ip and not ip.startswith("fe80::") and ip != "::1":
return ip.split("%")[0]
except (ValueError, OSError) as e:
logger.warning(
f"{e} Can not get local ip from NIC. Please verify whether SGLANG_LOCAL_IP_NIC is set correctly."
)
return None
def get_local_ip_by_remote() -> Optional[str]:
# Google's public DNS servers, used to discover the local IP.
# UDP connect doesn't send packets; it just selects the right source address.
# https://developers.google.com/speed/public-dns/docs/using#addresses
# Try IPv4 first, then IPv6. getaddrinfo on a literal IP returns exactly
# one result, so we unpack directly instead of looping.
for dns_host, dns_port in [("8.8.8.8", 80), ("2001:4860:4860::8888", 80)]:
try:
family, socktype, proto, _, sockaddr = socket.getaddrinfo(
dns_host,
dns_port,
socket.AF_UNSPEC,
socket.SOCK_DGRAM,
0,
socket.AI_ADDRCONFIG,
)[0]
with socket.socket(family, socktype, proto) as s:
s.connect(sockaddr)
return s.getsockname()[0]
except (socket.gaierror, OSError):
continue
# Fallback: resolve the local hostname to an IP address via /etc/hosts or DNS.
# Unreliable — many machines resolve hostname to 127.0.0.1, so we skip loopback.
try:
hostname = socket.gethostname()
ip = socket.getaddrinfo(
hostname, None, socket.AF_UNSPEC, 0, 0, socket.AI_ADDRCONFIG
)[0][4][0]
if ip and ip not in ("127.0.0.1", "0.0.0.0", "::1"):
return ip
except Exception:
pass
logger.warning("Can not get local ip by remote")
return None
def get_local_ip_auto(fallback: str = None) -> str:
"""
Automatically detect the local IP address using multiple fallback strategies.
This function attempts to obtain the local IP address through several methods.
If all methods fail, it returns the specified fallback value or raises an exception.
Args:
fallback (str, optional): Fallback IP address to return if all detection
methods fail. For server applications, explicitly set this to
"0.0.0.0" (IPv4) or "::" (IPv6) to bind to all available interfaces.
Defaults to None.
Returns:
str: The detected local IP address, or the fallback value if detection fails.
Raises:
ValueError: If IP detection fails and no fallback value is provided.
Note:
The function tries detection methods in the following order:
1. Direct IP detection via get_ip()
2. Network interface enumeration via get_local_ip_by_nic()
3. Remote connection method via get_local_ip_by_remote()
"""
# Try environment variable
host_ip = os.getenv("SGLANG_HOST_IP", "") or os.getenv("HOST_IP", "")
if host_ip:
return host_ip
logger.debug("get_ip failed")
# Fallback
if ip := get_local_ip_by_nic():
return ip
logger.debug("get_local_ip_by_nic failed")
# Fallback
if ip := get_local_ip_by_remote():
return ip
logger.debug("get_local_ip_by_remote failed")
if fallback:
return fallback
raise ValueError("Can not get local ip")
# TODO(hebiao064): Accelerate FA3 Spec Decode with topk > 1. # TODO(hebiao064): Accelerate FA3 Spec Decode with topk > 1.
# TODO(hebiao064): Improve the acc rate for FA3 Spec Decode with topk == 1 and page_size > 1. # TODO(hebiao064): Improve the acc rate for FA3 Spec Decode with topk == 1 and page_size > 1.
def is_no_spec_infer_or_topk_one(server_args): def is_no_spec_infer_or_topk_one(server_args):
+382 -1
View File
@@ -1,9 +1,390 @@
from __future__ import annotations from __future__ import annotations
import ipaddress import ipaddress
import logging
import os
import socket import socket
import time
from dataclasses import dataclass from dataclasses import dataclass
from typing import Tuple from typing import Optional, Tuple, Union
import psutil
import zmq
logger = logging.getLogger(__name__)
def get_open_port() -> int:
port = os.getenv("SGLANG_PORT")
if port is not None:
port = int(port)
while True:
if is_port_available(port):
return port
logger.info("Port %d is already in use, trying port %d", port, port + 1)
port += 1
sock = try_bind_socket()
port = sock.getsockname()[1]
sock.close()
return port
def is_valid_ipv6_address(address: str) -> bool:
try:
ipaddress.IPv6Address(address)
return True
except ValueError:
return False
def find_process_using_port(port: int) -> Optional[psutil.Process]:
for conn in psutil.net_connections(kind="inet"):
if conn.laddr.port == port:
try:
return psutil.Process(conn.pid)
except psutil.NoSuchProcess:
# It could happen by race condition (the proc dies when psutil.Process is called).
pass
return None
def wait_port_available(
port: int, port_name: str, timeout_s: int = 30, raise_exception: bool = True
) -> bool:
for i in range(timeout_s):
if is_port_available(port):
return True
if i > 10 and i % 5 == 0:
process = find_process_using_port(port)
if process is None:
logger.warning(
f"The port {port} is in use, but we could not find the process that uses it."
)
pid = process.pid
error_message = f"{port_name} is used by a process already. {process.name()=}' {process.cmdline()=} {process.status()=} {pid=}"
logger.info(
f"port {port} is in use. Waiting for {i} seconds for {port_name} to be available. {error_message}"
)
time.sleep(0.1)
if raise_exception:
raise ValueError(
f"{port_name} at {port} is not available in {timeout_s} seconds. {error_message}"
)
return False
def _get_addrinfos_for_bind(host=None, port=0):
"""Return deduplicated addrinfo tuples for binding (one per address family).
Args:
host: Bind address. None (with AI_PASSIVE) resolves to wildcard
addresses (0.0.0.0 / ::) suitable for accepting on all interfaces.
port: Port number. 0 lets the OS assign an available ephemeral port.
Flags:
AI_ADDRCONFIG — only return families actually configured on this host.
AI_PASSIVE — return wildcard addresses suitable for bind().
Falls back to AF_INET if getaddrinfo fails (e.g. DNS misconfiguration).
"""
try:
infos = socket.getaddrinfo(
host,
port,
socket.AF_UNSPEC,
socket.SOCK_STREAM,
0,
socket.AI_ADDRCONFIG | socket.AI_PASSIVE,
)
deduped = []
seen_families = set()
for info in infos:
if info[0] not in seen_families:
seen_families.add(info[0])
deduped.append(info)
# Prefer IPv4 so that callers without an explicit host get consistent
# behaviour across platforms (some OSes list IPv6 first).
deduped.sort(key=lambda x: (x[0] != socket.AF_INET,))
return deduped
except socket.gaierror:
fallback_host = "0.0.0.0" if host is None else host
return [(socket.AF_INET, socket.SOCK_STREAM, 0, "", (fallback_host, port))]
def try_bind_socket(host=None, port=0, *, reuse_addr=True, listen=False):
"""Bind a TCP socket on the first available address family (IPv4/IPv6).
Iterates over address families returned by _get_addrinfos_for_bind and
returns the first socket that successfully binds.
Args:
host: Bind address. None binds to all interfaces (0.0.0.0 / ::).
port: Port number. 0 lets the OS assign an available ephemeral port;
use sock.getsockname()[1] to retrieve the assigned port.
reuse_addr: Set SO_REUSEADDR to allow quick port reuse after close.
listen: Call listen(1) after bind, making the socket ready to accept.
Returns:
The bound socket. Caller is responsible for closing it.
Raises:
OSError: If bind fails on all configured address families.
"""
for family, socktype, proto, _, sockaddr in _get_addrinfos_for_bind(host, port):
sock = socket.socket(family, socktype, proto)
try:
if family == socket.AF_INET6:
sock.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
if reuse_addr:
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
sock.bind(sockaddr)
if listen:
sock.listen(1)
return sock
except (OSError, OverflowError):
sock.close()
raise OSError(f"Could not bind port {port} on any configured address family")
def is_port_available(port):
"""Return whether a port is available on all configured address families."""
try:
for family, socktype, proto, _, sockaddr in _get_addrinfos_for_bind(port=port):
sock = socket.socket(family, socktype, proto)
try:
if family == socket.AF_INET6:
sock.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
sock.bind(sockaddr)
finally:
sock.close()
return True
except (OSError, OverflowError):
return False
def get_free_port():
sock = try_bind_socket()
port = sock.getsockname()[1]
sock.close()
return port
def bind_port(port):
"""Bind to a specific port, assuming it's available."""
return try_bind_socket(port=port, listen=True)
def get_zmq_socket_on_host(
context: zmq.Context,
socket_type: zmq.SocketType,
host: Optional[str] = None,
) -> Tuple[int, zmq.Socket]:
"""Create and configure a ZeroMQ socket.
Args:
context: ZeroMQ context to create the socket from.
socket_type: Type of ZeroMQ socket to create.
host: Optional host to bind/connect to, without "tcp://" prefix. If None, binds to "tcp://*".
Returns:
Tuple of (port, socket) where port is the randomly assigned TCP port.
"""
socket = context.socket(socket_type)
# Bind to random TCP port, auto-wrapping IPv6 and setting zmq.IPV6 flag
config_socket(socket, socket_type)
if host:
if is_valid_ipv6_address(host):
socket.setsockopt(zmq.IPV6, 1)
bind_host = f"tcp://[{host}]"
else:
bind_host = f"tcp://{host}"
else:
bind_host = "tcp://*"
port = socket.bind_to_random_port(bind_host)
return port, socket
def config_socket(socket, socket_type: zmq.SocketType):
mem = psutil.virtual_memory()
total_mem = mem.total / 1024**3
available_mem = mem.available / 1024**3
if total_mem > 32 and available_mem > 16:
buf_size = int(0.5 * 1024**3)
else:
buf_size = -1
def set_send_opt():
socket.setsockopt(zmq.SNDHWM, 0)
socket.setsockopt(zmq.SNDBUF, buf_size)
def set_recv_opt():
socket.setsockopt(zmq.RCVHWM, 0)
socket.setsockopt(zmq.RCVBUF, buf_size)
if socket_type == zmq.PUSH:
set_send_opt()
elif socket_type == zmq.PULL:
set_recv_opt()
elif socket_type in [zmq.DEALER, zmq.REQ, zmq.REP]:
set_send_opt()
set_recv_opt()
else:
raise ValueError(f"Unsupported socket type: {socket_type}")
def get_local_ip_by_nic(interface: str = None) -> Optional[str]:
if not (interface := interface or os.environ.get("SGLANG_LOCAL_IP_NIC", None)):
return None
try:
import netifaces
except ImportError as e:
raise ImportError(
"Environment variable SGLANG_LOCAL_IP_NIC requires package netifaces, please install it through 'pip install netifaces'"
) from e
try:
addresses = netifaces.ifaddresses(interface)
if netifaces.AF_INET in addresses:
for addr_info in addresses[netifaces.AF_INET]:
ip = addr_info.get("addr")
if ip and ip != "127.0.0.1" and ip != "0.0.0.0":
return ip
if netifaces.AF_INET6 in addresses:
for addr_info in addresses[netifaces.AF_INET6]:
ip = addr_info.get("addr")
if ip and not ip.startswith("fe80::") and ip != "::1":
return ip.split("%")[0]
except (ValueError, OSError) as e:
logger.warning(
f"{e} Can not get local ip from NIC. Please verify whether SGLANG_LOCAL_IP_NIC is set correctly."
)
return None
def get_local_ip_by_remote() -> Optional[str]:
# Google's public DNS servers, used to discover the local IP.
# UDP connect doesn't send packets; it just selects the right source address.
# https://developers.google.com/speed/public-dns/docs/using#addresses
# Try IPv4 first, then IPv6. getaddrinfo on a literal IP returns exactly
# one result, so we unpack directly instead of looping.
for dns_host, dns_port in [("8.8.8.8", 80), ("2001:4860:4860::8888", 80)]:
try:
family, socktype, proto, _, sockaddr = socket.getaddrinfo(
dns_host,
dns_port,
socket.AF_UNSPEC,
socket.SOCK_DGRAM,
0,
socket.AI_ADDRCONFIG,
)[0]
with socket.socket(family, socktype, proto) as s:
s.connect(sockaddr)
return s.getsockname()[0]
except (socket.gaierror, OSError):
continue
# Fallback: resolve the local hostname to an IP address via /etc/hosts or DNS.
# Unreliable — many machines resolve hostname to 127.0.0.1, so we skip loopback.
try:
hostname = socket.gethostname()
ip = socket.getaddrinfo(
hostname, None, socket.AF_UNSPEC, 0, 0, socket.AI_ADDRCONFIG
)[0][4][0]
if ip and ip not in ("127.0.0.1", "0.0.0.0", "::1"):
return ip
except Exception:
pass
logger.warning("Can not get local ip by remote")
return None
def get_local_ip_auto(fallback: str = None) -> str:
"""
Automatically detect the local IP address using multiple fallback strategies.
This function attempts to obtain the local IP address through several methods.
If all methods fail, it returns the specified fallback value or raises an exception.
Args:
fallback (str, optional): Fallback IP address to return if all detection
methods fail. For server applications, explicitly set this to
"0.0.0.0" (IPv4) or "::" (IPv6) to bind to all available interfaces.
Defaults to None.
Returns:
str: The detected local IP address, or the fallback value if detection fails.
Raises:
ValueError: If IP detection fails and no fallback value is provided.
Note:
The function tries detection methods in the following order:
1. Direct IP detection via get_ip()
2. Network interface enumeration via get_local_ip_by_nic()
3. Remote connection method via get_local_ip_by_remote()
"""
# Try environment variable
host_ip = os.getenv("SGLANG_HOST_IP", "") or os.getenv("HOST_IP", "")
if host_ip:
return host_ip
logger.debug("get_ip failed")
# Fallback
if ip := get_local_ip_by_nic():
return ip
logger.debug("get_local_ip_by_nic failed")
# Fallback
if ip := get_local_ip_by_remote():
return ip
logger.debug("get_local_ip_by_remote failed")
if fallback:
return fallback
raise ValueError("Can not get local ip")
def get_zmq_socket(
context: zmq.Context,
socket_type: zmq.SocketType,
endpoint: Optional[str] = None,
bind: bool = True,
) -> Union[zmq.Socket, Tuple[int, zmq.Socket]]:
"""Create and configure a ZeroMQ socket.
Args:
context: ZeroMQ context to create the socket from.
socket_type: Type of ZeroMQ socket to create.
endpoint: Optional endpoint to bind/connect to. If None, binds to a random TCP port.
bind: Whether to bind (True) or connect (False) to the endpoint. Ignored if endpoint is None.
Returns:
If endpoint is None: Tuple of (port, socket) where port is the randomly assigned TCP port.
If endpoint is provided: The configured ZeroMQ socket.
"""
socket = context.socket(socket_type)
if endpoint is None:
# Bind to random TCP port
config_socket(socket, socket_type)
port = socket.bind_to_random_port("tcp://*")
return port, socket
else:
# Handle IPv6 if endpoint contains brackets
if endpoint.find("[") != -1:
socket.setsockopt(zmq.IPV6, 1)
config_socket(socket, socket_type)
if bind:
socket.bind(endpoint)
else:
socket.connect(endpoint)
return socket
def _is_ipv6(host: str) -> bool: def _is_ipv6(host: str) -> bool:
+1 -1
View File
@@ -38,11 +38,11 @@ from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
get_device, get_device,
is_cuda, is_cuda,
is_port_available,
is_xpu, is_xpu,
kill_process_tree, kill_process_tree,
retry, retry,
) )
from sglang.srt.utils.network import is_port_available
from sglang.test.run_eval import run_eval from sglang.test.run_eval import run_eval
from sglang.utils import get_exception_traceback, normalize_base_url from sglang.utils import get_exception_traceback, normalize_base_url
+1 -1
View File
@@ -401,7 +401,7 @@ def reserve_port(host, start=30000, end=40000):
Reserve an available port by trying to bind a socket. Reserve an available port by trying to bind a socket.
Returns a tuple (port, lock_socket) where `lock_socket` is kept open to hold the lock. Returns a tuple (port, lock_socket) where `lock_socket` is kept open to hold the lock.
""" """
from sglang.srt.utils.common import try_bind_socket from sglang.srt.utils.network import try_bind_socket
candidates = list(range(start, end)) candidates = list(range(start, end))
random.shuffle(candidates) random.shuffle(candidates)
@@ -15,7 +15,7 @@ from setproctitle import setproctitle
from sglang_router.launch_router import RouterArgs, launch_router from sglang_router.launch_router import RouterArgs, launch_router
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_port_available from sglang.srt.utils.network import is_port_available
def setup_logger(): def setup_logger():
+2 -1
View File
@@ -12,7 +12,8 @@ import zmq
from sglang import Engine from sglang import Engine
from sglang.srt.observability.trace import * from sglang.srt.observability.trace import *
from sglang.srt.observability.trace import get_cur_time_ns, set_global_trace_level from sglang.srt.observability.trace import get_cur_time_ns, set_global_trace_level
from sglang.srt.utils import get_zmq_socket, kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.network import get_zmq_socket
from sglang.test.test_utils import ( from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
@@ -8,7 +8,8 @@ import grpc
import zmq import zmq
from grpc_health.v1 import health_pb2, health_pb2_grpc from grpc_health.v1 import health_pb2, health_pb2_grpc
from sglang.srt.utils import get_zmq_socket_on_host, kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.srt.utils.network import get_zmq_socket_on_host
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.mmmu_vlm_kit import _run_lmms_eval_with_retry from sglang.test.kits.mmmu_vlm_kit import _run_lmms_eval_with_retry
from sglang.test.server_fixtures.disaggregation_fixture import ( from sglang.test.server_fixtures.disaggregation_fixture import (
+4 -4
View File
@@ -3,7 +3,7 @@ import socket
import unittest import unittest
from unittest.mock import patch from unittest.mock import patch
from sglang.srt.utils.common import ( from sglang.srt.utils.network import (
_get_addrinfos_for_bind, _get_addrinfos_for_bind,
bind_port, bind_port,
get_free_port, get_free_port,
@@ -78,7 +78,7 @@ class TestTryBindSocket(CustomTestCase):
def test_gaierror_fallback(self): def test_gaierror_fallback(self):
"""_get_addrinfos_for_bind should fall back to AF_INET on gaierror.""" """_get_addrinfos_for_bind should fall back to AF_INET on gaierror."""
with patch( with patch(
"sglang.srt.utils.common.socket.getaddrinfo", "sglang.srt.utils.network.socket.getaddrinfo",
side_effect=socket.gaierror("mocked"), side_effect=socket.gaierror("mocked"),
): ):
infos = _get_addrinfos_for_bind() infos = _get_addrinfos_for_bind()
@@ -90,7 +90,7 @@ class TestTryBindSocket(CustomTestCase):
def test_gaierror_fallback_preserves_host(self): def test_gaierror_fallback_preserves_host(self):
"""Fallback should use the provided host, not default to 0.0.0.0.""" """Fallback should use the provided host, not default to 0.0.0.0."""
with patch( with patch(
"sglang.srt.utils.common.socket.getaddrinfo", "sglang.srt.utils.network.socket.getaddrinfo",
side_effect=socket.gaierror("mocked"), side_effect=socket.gaierror("mocked"),
): ):
infos = _get_addrinfos_for_bind(host="10.0.0.1", port=8080) infos = _get_addrinfos_for_bind(host="10.0.0.1", port=8080)
@@ -187,7 +187,7 @@ class TestReservePort(CustomTestCase):
def test_reserve_port_no_free_port_raises(self): def test_reserve_port_no_free_port_raises(self):
"""reserve_port should raise RuntimeError if no port is available.""" """reserve_port should raise RuntimeError if no port is available."""
with patch( with patch(
"sglang.srt.utils.common.try_bind_socket", "sglang.srt.utils.network.try_bind_socket",
side_effect=OSError("mocked"), side_effect=OSError("mocked"),
): ):
with self.assertRaises(RuntimeError): with self.assertRaises(RuntimeError):