[PD] Un-blacklist mooncake sessions when probe succeeds (#25287)

This commit is contained in:
Kevin Flansburg
2026-05-20 11:16:28 +08:00
committed by GitHub
parent 5fe655bf26
commit 3b2178c412
3 changed files with 61 additions and 0 deletions
@@ -13,6 +13,7 @@ from typing import List, Optional, Tuple
import numpy as np
import numpy.typing as npt
from prometheus_client import Counter
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll, StateType
from sglang.srt.disaggregation.common.conn import (
@@ -47,6 +48,11 @@ from sglang.srt.utils.network import NetworkAddress
logger = logging.getLogger(__name__)
FAILED_SESSION_RECOVERIES = Counter(
"sglang:failed_session_recoveries_total",
"Number of mooncake_session_ids un-blacklisted via probe.",
)
class KVTransferError(Exception):
def __init__(self, bootstrap_room: int, failure_reason: str):
@@ -238,6 +244,19 @@ class MooncakeKVManager(CommonKVManager):
),
daemon=True,
).start()
self.enable_failed_session_probe = (
envs.SGLANG_ENABLE_FAILED_SESSION_PROBE.get()
)
if self.enable_failed_session_probe:
self.failed_session_probe_interval = (
envs.SGLANG_FAILED_SESSION_PROBE_INTERVAL_S.get()
)
self._failed_session_probe_shutdown = threading.Event()
threading.Thread(
target=self._failed_session_probe_loop,
name="MooncakeFailedSessionProbe",
daemon=True,
).start()
elif self.disaggregation_mode == DisaggregationMode.DECODE:
self._staging_ctx = DecodeStagingContext() if self.enable_staging else None
if self.enable_staging:
@@ -1564,6 +1583,43 @@ class MooncakeKVManager(CommonKVManager):
def get_session_id(self):
return self.engine.get_session_id()
def _run_one_probe_pass(self) -> None:
with self.session_lock:
snapshot = list(self.failed_sessions)
for session_id in snapshot:
send_probe = getattr(self.engine, "send_probe", None)
if send_probe is None:
rc = -1
else:
try:
rc = send_probe(session_id)
except Exception as e:
logger.warning("send_probe(%s) raised: %s", session_id, e)
continue
if rc == 0:
with self.session_lock:
was_blacklisted = session_id in self.failed_sessions
self.failed_sessions.discard(session_id)
self.session_failures.pop(session_id, None)
if was_blacklisted:
logger.info(
"Session %s recovered via probe; un-blacklisted",
session_id,
)
FAILED_SESSION_RECOVERIES.inc()
else:
logger.debug("Probe still failing for %s (rc=%d)", session_id, rc)
def _failed_session_probe_loop(self) -> None:
logger.info(
"Starting failed-session probe loop (interval=%.1fs)",
self.failed_session_probe_interval,
)
while not self._failed_session_probe_shutdown.wait(
self.failed_session_probe_interval
):
self._run_one_probe_pass()
def _handle_node_failure(self, failed_bootstrap_addr):
with self.connection_lock:
keys_to_remove = [
@@ -254,6 +254,9 @@ class MooncakeTransferEngine:
def get_session_id(self):
return self.session_id
def send_probe(self, peer_session_id: str) -> int:
return self.engine.send_probe(peer_session_id)
def get_engine(self):
return self.engine.get_engine()
+2
View File
@@ -306,6 +306,8 @@ class Envs:
ENABLE_ASCEND_TRANSFER_WITH_MOONCAKE = EnvBool(False)
ASCEND_NPU_PHY_ID = EnvInt(-1)
SGLANG_MOONCAKE_SEND_AUX_TCP = EnvBool(False)
SGLANG_ENABLE_FAILED_SESSION_PROBE = EnvBool(False)
SGLANG_FAILED_SESSION_PROBE_INTERVAL_S = EnvFloat(30.0)
# Mooncake Store
SGLANG_HICACHE_MOONCAKE_CONFIG_PATH = EnvStr(None)