From 6a012fbb2dc7a36b12013860e6dab47a97e369d4 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Thu, 11 Jun 2026 19:51:12 +0800 Subject: [PATCH] [PD] Fix ZMQ stale socket reconnection in PD disaggregation (#27796) Signed-off-by: Shangming Cai Co-authored-by: Abatom --- .../sglang/srt/disaggregation/common/conn.py | 25 +++++++++++++++ .../srt/disaggregation/mooncake/conn.py | 31 ------------------- 2 files changed, 25 insertions(+), 31 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 087c90320..67115dfd2 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -710,6 +710,15 @@ class CommonKVManager(BaseKVManager): keys_to_remove = [ k for k in self.connection_pool if k.startswith(failed_bootstrap_addr) ] + # Collect TCP endpoints from cached bootstrap_infos before deletion + stale_endpoints = set() + for k in keys_to_remove: + for info in self.connection_pool[k]: + ip = info.get("rank_ip") + port = info.get("rank_port") + if ip and port: + na = NetworkAddress(ip, int(port)) + stale_endpoints.add(na.to_tcp()) for k in keys_to_remove: del self.connection_pool[k] self.prefill_info_table.pop(failed_bootstrap_addr, None) @@ -719,6 +728,9 @@ class CommonKVManager(BaseKVManager): ) self.addr_to_rooms_tracker.pop(failed_bootstrap_addr, None) + for endpoint in stale_endpoints: + CommonKVReceiver.disconnect_endpoint(endpoint) + affected_rooms = [] for room in possible_affected_rooms: if ( @@ -1082,6 +1094,19 @@ class CommonKVReceiver(BaseKVReceiver): cls._socket_locks[endpoint] = threading.Lock() return cls._socket_cache[endpoint], cls._socket_locks[endpoint] + @classmethod + def disconnect_endpoint(cls, endpoint: str): + with cls._global_lock: + sock = cls._socket_cache.pop(endpoint, None) + lock = cls._socket_locks.pop(endpoint, None) + if sock: + if lock: + with lock: + sock.close() + else: + sock.close() + logger.debug(f"Disconnected stale ZMQ PUSH socket (receiver): {endpoint}") + @classmethod def _connect_to_bootstrap_server(cls, bootstrap_info: dict): ip_address = bootstrap_info["rank_ip"] diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index e2b7e0a9f..ffc6c6cfd 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -1651,37 +1651,6 @@ class MooncakeKVManager(CommonKVManager): ): self._run_one_probe_pass() - def _handle_node_failure(self, failed_bootstrap_addr): - with self.connection_lock: - keys_to_remove = [ - k for k in self.connection_pool if k.startswith(failed_bootstrap_addr) - ] - for k in keys_to_remove: - del self.connection_pool[k] - - possible_affected_rooms = self.addr_to_rooms_tracker.get( - failed_bootstrap_addr, [] - ) - self.prefill_info_table.pop(failed_bootstrap_addr, None) - self.addr_to_rooms_tracker.pop(failed_bootstrap_addr, None) - - # Report the requests associated with the failed bootstrap addr and mark their status as KVPoll.Failed - affected_rooms = [] - for room in possible_affected_rooms: - if ( - room in self.request_status - and self.check_status(room) != KVPoll.Success - ): - self.record_failure( - room, - f"Losing connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr})", - ) - self.update_status(room, KVPoll.Failed) - affected_rooms.append(room) - logger.error( - f"Losing connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr}), {len(affected_rooms)} requests affected" - ) - class MooncakeKVSender(CommonKVSender):