From be45745f3858f12f4ad006adf9986bf7b0dcd7be Mon Sep 17 00:00:00 2001 From: Zhonghua Deng Date: Thu, 11 Jun 2026 19:32:31 +0800 Subject: [PATCH] [EPD] fix: zmq PUSH socket reconnect-aware connection management with tcp keepalive (#27039) --- .../sglang/srt/disaggregation/common/conn.py | 51 +++++++++++++++---- .../srt/disaggregation/encode_server.py | 2 +- 2 files changed, 43 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index e88024c5a..087c90320 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -6,7 +6,6 @@ import logging import threading import time from collections import defaultdict -from functools import cache from typing import Dict, List, Optional, Set, Tuple, Union import numpy as np @@ -144,13 +143,16 @@ class CommonKVManager(BaseKVManager): ) # bind zmq socket - context = zmq.Context() + self._zmq_ctx = zmq.Context() self.rank_port, self.server_socket = get_zmq_socket_on_host( - context, zmq.PULL, host=self.local_ip + self._zmq_ctx, zmq.PULL, host=self.local_ip ) logger.debug(f"kv manager bind to {self.local_ip}:{self.rank_port}") self.request_status: Dict[int, KVPoll] = {} + self._socket_cache: Dict[str, zmq.Socket] = {} + self._monitor_cache: Dict[str, zmq.Socket] = {} + self._socket_lock = threading.Lock() self.failure_records: Dict[int, str] = {} self.failure_lock = threading.Lock() @@ -446,13 +448,44 @@ class CommonKVManager(BaseKVManager): f"Prefill instance failed to register to bootstrap server after {max_retries} retries" ) - @cache def _connect(self, endpoint: str, is_ipv6: bool = False): - socket = zmq.Context().socket(zmq.PUSH) - if is_ipv6: - socket.setsockopt(zmq.IPV6, 1) - socket.connect(endpoint) - return socket + with self._socket_lock: + sock = self._socket_cache.get(endpoint) + if sock is not None: + monitor = self._monitor_cache.get(endpoint) + disconnected = False + if monitor is not None: + try: + monitor.recv_multipart(zmq.NOBLOCK) + disconnected = True + except zmq.Again: + pass + except zmq.ZMQError: + disconnected = True + if not disconnected: + return sock + sock.close(linger=0) + if monitor is not None: + monitor.close() + self._socket_cache.pop(endpoint, None) + self._monitor_cache.pop(endpoint, None) + + sock = self._zmq_ctx.socket(zmq.PUSH) + if is_ipv6: + sock.setsockopt(zmq.IPV6, 1) + sock.setsockopt(zmq.RECONNECT_IVL, -1) + sock.setsockopt(zmq.SNDTIMEO, 30000) + sock.setsockopt(zmq.LINGER, 0) + sock.setsockopt(zmq.TCP_KEEPALIVE, 1) + sock.setsockopt(zmq.TCP_KEEPALIVE_IDLE, 30) + sock.setsockopt(zmq.TCP_KEEPALIVE_INTVL, 5) + sock.setsockopt(zmq.TCP_KEEPALIVE_CNT, 3) + sock.connect(endpoint) + self._socket_cache[endpoint] = sock + self._monitor_cache[endpoint] = sock.get_monitor_socket( + zmq.EVENT_DISCONNECTED + ) + return sock def get_mha_kv_ptrs_with_pp( self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int] diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 7f5bce2b7..54fbea330 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -1639,7 +1639,7 @@ class MMEncoder: else: sock.send_multipart([serialized_data], copy=False) finally: - sock.close() + sock.close(linger=5000) await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)