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)