diff --git a/docs_new/docs/references/environment_variables.mdx b/docs_new/docs/references/environment_variables.mdx
index 4a847ea47..67750563c 100644
--- a/docs_new/docs/references/environment_variables.mdx
+++ b/docs_new/docs/references/environment_variables.mdx
@@ -1428,6 +1428,11 @@ SGLang supports various environment variables that can be used to configure its
Consecutive heartbeat failures tolerated before a peer is considered dead. |
2 |
+
+ SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT |
+ Send timeout (seconds) for decode's ZMQ sockets to prefill peers. Raise if healthy sends exceed the default. |
+ 1 |
+
SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL |
Interval (seconds) for cleaning up stale bootstrap entries. |
diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py
index 064165f46..402c930aa 100644
--- a/python/sglang/srt/disaggregation/common/conn.py
+++ b/python/sglang/srt/disaggregation/common/conn.py
@@ -1260,10 +1260,13 @@ class CommonKVReceiver(BaseKVReceiver):
return
self.bootstrap_infos = bootstrap_infos
- self.kv_mgr.connection_pool[bootstrap_key] = self.bootstrap_infos
- # Register kv_args only once to prefill KVManager according to the info fetched from the bootstrap server
- self._register_kv_args()
+ # Register kv_args only once to prefill KVManager according to the info fetched
+ # from the bootstrap server. Do this before caching in connection_pool so a failed
+ # registration does not leave a stale entry that later requests would reuse.
+ if not self._register_kv_args():
+ return
+ self.kv_mgr.connection_pool[bootstrap_key] = self.bootstrap_infos
else:
self.bootstrap_infos = self.kv_mgr.connection_pool[bootstrap_key]
@@ -1322,6 +1325,11 @@ class CommonKVReceiver(BaseKVReceiver):
if is_ipv6:
sock.setsockopt(zmq.IPV6, 1)
sock.setsockopt(zmq.LINGER, 0)
+ # Bound send so a dead peer cannot block the scheduler forever.
+ sock.setsockopt(
+ zmq.SNDTIMEO,
+ envs.SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT.get() * 1000,
+ )
sock.connect(endpoint)
cls._socket_cache[endpoint] = sock
cls._socket_locks[endpoint] = threading.Lock()
@@ -1348,8 +1356,8 @@ class CommonKVReceiver(BaseKVReceiver):
sock, lock = cls._connect(na.to_tcp(), is_ipv6=na.is_ipv6)
return sock, lock
- def _register_kv_args(self):
- pass
+ def _register_kv_args(self) -> bool:
+ return True
def send_metadata(
self,
diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py
index 9f04f170f..182844ce5 100644
--- a/python/sglang/srt/disaggregation/mooncake/conn.py
+++ b/python/sglang/srt/disaggregation/mooncake/conn.py
@@ -12,6 +12,7 @@ from typing import List, Optional, Tuple, Union
import numpy as np
import numpy.typing as npt
+import zmq
from prometheus_client import Counter
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType
@@ -1901,7 +1902,7 @@ class MooncakeKVReceiver(CommonKVReceiver):
self.init_time = None
super().__init__(mgr, bootstrap_addr, bootstrap_room)
- def _register_kv_args(self):
+ def _register_kv_args(self) -> bool:
for bootstrap_info in self.bootstrap_infos:
packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
@@ -1936,25 +1937,35 @@ class MooncakeKVReceiver(CommonKVReceiver):
staging_total_size_str = b""
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
- with lock:
- sock.send_multipart(
- [
- "None".encode("ascii"),
- self.kv_mgr.local_ip.encode("ascii"),
- str(self.kv_mgr.rank_port).encode("ascii"),
- self.session_id.encode("ascii"),
- packed_kv_data_ptrs,
- packed_aux_data_ptrs,
- packed_state_data_ptrs,
- dst_tp_rank,
- dst_attn_tp_size,
- dst_kv_item_len,
- packed_state_item_lens,
- packed_state_dim_per_tensor,
- packed_staging_base_ptr,
- staging_total_size_str,
- ]
+ try:
+ with lock:
+ sock.send_multipart(
+ [
+ "None".encode("ascii"),
+ self.kv_mgr.local_ip.encode("ascii"),
+ str(self.kv_mgr.rank_port).encode("ascii"),
+ self.session_id.encode("ascii"),
+ packed_kv_data_ptrs,
+ packed_aux_data_ptrs,
+ packed_state_data_ptrs,
+ dst_tp_rank,
+ dst_attn_tp_size,
+ dst_kv_item_len,
+ packed_state_item_lens,
+ packed_state_dim_per_tensor,
+ packed_staging_base_ptr,
+ staging_total_size_str,
+ ]
+ )
+ except zmq.ZMQError:
+ self.kv_mgr.record_failure(
+ self.bootstrap_room,
+ f"_register_kv_args to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
+ self.conclude_state = KVPoll.Failed
+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
+ return False
+ return True
def send_metadata(
self,
@@ -1983,25 +1994,33 @@ class MooncakeKVReceiver(CommonKVReceiver):
for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
is_dummy = bootstrap_info["is_dummy"]
-
- with lock:
- sock.send_multipart(
- [
- str(self.bootstrap_room).encode("ascii"),
- self.kv_mgr.local_ip.encode("ascii"),
- str(self.kv_mgr.rank_port).encode("ascii"),
- self.session_id.encode("ascii"),
- kv_indices.tobytes() if not is_dummy else b"",
- str(aux_index).encode("ascii") if not is_dummy else b"",
- (
- pack_int_lists(state_indices, "i")
- if not is_dummy and state_indices
- else b""
- ),
- str(self.required_dst_info_num).encode("ascii"),
- str(decode_prefix_len or 0).encode("ascii"),
- ]
+ try:
+ with lock:
+ sock.send_multipart(
+ [
+ str(self.bootstrap_room).encode("ascii"),
+ self.kv_mgr.local_ip.encode("ascii"),
+ str(self.kv_mgr.rank_port).encode("ascii"),
+ self.session_id.encode("ascii"),
+ kv_indices.tobytes() if not is_dummy else b"",
+ str(aux_index).encode("ascii") if not is_dummy else b"",
+ (
+ pack_int_lists(state_indices, "i")
+ if not is_dummy and state_indices
+ else b""
+ ),
+ str(self.required_dst_info_num).encode("ascii"),
+ str(decode_prefix_len or 0).encode("ascii"),
+ ]
+ )
+ except zmq.ZMQError:
+ self.kv_mgr.record_failure(
+ self.bootstrap_room,
+ f"send_metadata to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
+ self.conclude_state = KVPoll.Failed
+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
+ return
self.init_time = time.time()
def poll(self) -> KVPoll:
diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py
index b175ca541..ac591179f 100644
--- a/python/sglang/srt/disaggregation/mori/conn.py
+++ b/python/sglang/srt/disaggregation/mori/conn.py
@@ -1658,9 +1658,9 @@ class MoriKVReceiver(CommonKVReceiver):
return
self.kv_mgr.room_to_bootstrap_addr[self.bootstrap_room] = self.bootstrap_addr
- def _register_kv_args(self):
+ def _register_kv_args(self) -> bool:
if self.bootstrap_infos is None:
- return
+ return False
engine_desc_blob = self.kv_mgr.engine_desc.pack()
packed_kv_descs = _pack_mem_desc_list(self.kv_mgr.kv_mem_descs)
packed_aux_descs = _pack_mem_desc_list(self.kv_mgr.aux_mem_descs)
@@ -1678,25 +1678,35 @@ class MoriKVReceiver(CommonKVReceiver):
for bootstrap_info in self.bootstrap_infos:
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
- with lock:
- sock.send_multipart(
- [
- MORI_GUARD,
- "None".encode("ascii"),
- self.kv_mgr.local_ip.encode("ascii"),
- str(self.kv_mgr.rank_port).encode("ascii"),
- engine_desc_blob,
- packed_kv_descs,
- packed_aux_descs,
- packed_state_descs,
- gpu_id,
- decode_tp_size,
- decode_tp_rank,
- kv_item_len,
- packed_state_item_lens,
- packed_state_dim_per_tensor,
- ]
+ try:
+ with lock:
+ sock.send_multipart(
+ [
+ MORI_GUARD,
+ "None".encode("ascii"),
+ self.kv_mgr.local_ip.encode("ascii"),
+ str(self.kv_mgr.rank_port).encode("ascii"),
+ engine_desc_blob,
+ packed_kv_descs,
+ packed_aux_descs,
+ packed_state_descs,
+ gpu_id,
+ decode_tp_size,
+ decode_tp_rank,
+ kv_item_len,
+ packed_state_item_lens,
+ packed_state_dim_per_tensor,
+ ]
+ )
+ except zmq.ZMQError:
+ self.kv_mgr.record_failure(
+ self.bootstrap_room,
+ f"_register_kv_args to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
+ self.conclude_state = KVPoll.Failed
+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
+ return False
+ return True
def send_metadata(
self,
@@ -1727,21 +1737,30 @@ class MoriKVReceiver(CommonKVReceiver):
state_bytes = _pack_state_indices(normalized_state)
else:
state_bytes = b""
- with lock:
- sock.send_multipart(
- [
- MORI_GUARD,
- str(self.bootstrap_room).encode("ascii"),
- self.kv_mgr.local_ip.encode("ascii"),
- str(self.kv_mgr.rank_port).encode("ascii"),
- self.kv_mgr.engine_desc.key.encode("ascii"),
- kv_indices_bytes if not is_dummy else b"",
- aux_bytes if not is_dummy else b"",
- state_bytes,
- str(self.required_dst_info_num).encode("ascii"),
- decode_prefix_bytes,
- ]
+ try:
+ with lock:
+ sock.send_multipart(
+ [
+ MORI_GUARD,
+ str(self.bootstrap_room).encode("ascii"),
+ self.kv_mgr.local_ip.encode("ascii"),
+ str(self.kv_mgr.rank_port).encode("ascii"),
+ self.kv_mgr.engine_desc.key.encode("ascii"),
+ kv_indices_bytes if not is_dummy else b"",
+ aux_bytes if not is_dummy else b"",
+ state_bytes,
+ str(self.required_dst_info_num).encode("ascii"),
+ decode_prefix_bytes,
+ ]
+ )
+ except zmq.ZMQError:
+ self.kv_mgr.record_failure(
+ self.bootstrap_room,
+ f"send_metadata to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
+ self.conclude_state = KVPoll.Failed
+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
+ return
self.init_time = time.time()
def poll(self) -> KVPoll:
diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py
index 85660c5c7..0af144c4f 100644
--- a/python/sglang/srt/disaggregation/nixl/conn.py
+++ b/python/sglang/srt/disaggregation/nixl/conn.py
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
import numpy as np
import numpy.typing as npt
+import zmq
if TYPE_CHECKING:
from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo
@@ -2637,21 +2638,30 @@ class NixlKVReceiver(CommonKVReceiver):
if not is_dummy and state_indices is not None
else b""
)
- with lock:
- sock.send_multipart(
- [
- GUARD,
- str(self.bootstrap_room).encode("ascii"),
- self.kv_mgr.local_ip.encode("ascii"),
- str(self.kv_mgr.rank_port).encode("ascii"),
- self.kv_mgr.agent.name.encode("ascii"),
- kv_indices.tobytes() if not is_dummy else b"",
- str(aux_index).encode("ascii"),
- str(self.required_dst_info_num).encode("ascii"),
- packed_state_indices,
- str(decode_prefix_len or 0).encode("ascii"),
- ]
+ try:
+ with lock:
+ sock.send_multipart(
+ [
+ GUARD,
+ str(self.bootstrap_room).encode("ascii"),
+ self.kv_mgr.local_ip.encode("ascii"),
+ str(self.kv_mgr.rank_port).encode("ascii"),
+ self.kv_mgr.agent.name.encode("ascii"),
+ kv_indices.tobytes() if not is_dummy else b"",
+ str(aux_index).encode("ascii"),
+ str(self.required_dst_info_num).encode("ascii"),
+ packed_state_indices,
+ str(decode_prefix_len or 0).encode("ascii"),
+ ]
+ )
+ except zmq.ZMQError:
+ self.kv_mgr.record_failure(
+ self.bootstrap_room,
+ f"send_metadata to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
+ self.conclude_state = KVPoll.Failed
+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
+ return
# Mark that we expect state data if state_indices was provided.
# Match the prefill-side truthy check: an empty list means the
@@ -2687,9 +2697,8 @@ class NixlKVReceiver(CommonKVReceiver):
return self.conclude_state # type: ignore
return KVPoll.WaitingForInput # type: ignore
- def _register_kv_args(self):
+ def _register_kv_args(self) -> bool:
for bootstrap_info in self.bootstrap_infos:
- sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
packed_kv_data_ptrs = b"".join(
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
)
@@ -2729,31 +2738,42 @@ class NixlKVReceiver(CommonKVReceiver):
// self.kv_mgr.kv_args.kv_item_lens[0]
)
- with lock:
- sock.send_multipart(
- [
- GUARD,
- "None".encode("ascii"),
- self.kv_mgr.local_ip.encode("ascii"),
- str(self.kv_mgr.rank_port).encode("ascii"),
- self.kv_mgr.agent.name.encode("ascii"),
- self.kv_mgr.agent.get_agent_metadata(),
- packed_kv_data_ptrs,
- packed_aux_data_ptrs,
- packed_state_data_ptrs,
- str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
- str(self.kv_mgr.attn_tp_size).encode("ascii"),
- str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
- str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"),
- packed_state_item_lens,
- packed_state_dim_per_tensor,
- packed_staging_base_ptr,
- staging_total_size_str,
- str(dst_num_slots).encode("ascii"),
- packed_kv_data_mem_kinds,
- packed_kv_item_lens,
- ]
+ sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
+ try:
+ with lock:
+ sock.send_multipart(
+ [
+ GUARD,
+ "None".encode("ascii"),
+ self.kv_mgr.local_ip.encode("ascii"),
+ str(self.kv_mgr.rank_port).encode("ascii"),
+ self.kv_mgr.agent.name.encode("ascii"),
+ self.kv_mgr.agent.get_agent_metadata(),
+ packed_kv_data_ptrs,
+ packed_aux_data_ptrs,
+ packed_state_data_ptrs,
+ str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
+ str(self.kv_mgr.attn_tp_size).encode("ascii"),
+ str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
+ str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"),
+ packed_state_item_lens,
+ packed_state_dim_per_tensor,
+ packed_staging_base_ptr,
+ staging_total_size_str,
+ str(dst_num_slots).encode("ascii"),
+ packed_kv_data_mem_kinds,
+ packed_kv_item_lens,
+ ]
+ )
+ except zmq.ZMQError:
+ self.kv_mgr.record_failure(
+ self.bootstrap_room,
+ f"_register_kv_args to prefill {bootstrap_info.get('rank_ip')}:{bootstrap_info.get('rank_port')} failed",
)
+ self.conclude_state = KVPoll.Failed
+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
+ return False
+ return True
def failure_exception(self):
with self.kv_mgr.failure_lock:
diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py
index 5aea0593d..2753260cd 100644
--- a/python/sglang/srt/environ.py
+++ b/python/sglang/srt/environ.py
@@ -392,6 +392,7 @@ class Envs:
SGLANG_DISAGGREGATION_THREAD_POOL_SIZE = EnvInt(None)
SGLANG_DISAGGREGATION_QUEUE_SIZE = EnvInt(4)
SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT = EnvInt(300)
+ SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT = EnvInt(1)
SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0)
SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2)
SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300)