From 1e221720f082abf54d089c2f406bce54dee61f65 Mon Sep 17 00:00:00 2001 From: Ashfaq <105435085+Ashfaqbs@users.noreply.github.com> Date: Wed, 26 Aug 2026 08:19:05 +0530 Subject: [PATCH] fix(pd): guard remote_agents dict with a lock to fix NIXL concurrent-add race connect_add_remote_agent() did an unlocked check-then-add on self.remote_agents, and remove_remote_agent() did an unlocked pop. Both are called from several worker threads on the same NixlKVTransporter instance (recv/dispatch/accept_peer/request-page/ ready-page loops all run as separate threads sharing one transporter). When a peer briefly looks missing (e.g. after a transient NIXL_ERR_REMOTE_DISCONNECT removes it), two threads can both see it absent and both call nixl_agent.add_remote_agent() for the same peer concurrently, which can corrupt UCX endpoint state and abort the NIXL progress thread (observed as a ud_ep.c PSN assertion), taking down the router. Guard the whole check-then-add and check-then-remove sequences with a lock so only one thread can add or remove a given peer at a time. Fixes #1470 --- .../mode_backend/pd/nixl_kv_transporter.py | 47 +++++++++++-------- 1 file changed, 27 insertions(+), 20 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index bd5e11f05d..8bcd23abb5 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -1,6 +1,7 @@ import pickle import copy import os +import threading import time from dataclasses import dataclass from typing import Dict @@ -40,6 +41,10 @@ def __init__(self, node_id: int, tp_idx: int, kv_move_buffer: Tensor): self.nixl_agent = NixlWrapper(self.agent_name, conf) self._register_kv_move_buffer(kv_move_buffer=kv_move_buffer) self.remote_agents: Dict[str, PDAgentMetadata] = {} + # remote_agents is read and mutated from several worker threads (recv, dispatch, + # accept_peer, request/ready page loops, ...), so every access needs to go through + # this lock to avoid concurrent add/remove races on the same peer (see GH-1470). + self._remote_agents_lock = threading.Lock() return @property @@ -73,36 +78,38 @@ def _create_paged_xfer_handles(self, reg_desc: "nixlBind.nixlRegDList", page_num return self.nixl_agent.prep_xfer_dlist(agent_name, descs, "VRAM") def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): - if remote_agent.agent_name in self.remote_agents: - return + with self._remote_agents_lock: + if remote_agent.agent_name in self.remote_agents: + return - start_time = time.time() + start_time = time.time() - peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) - if isinstance(peer_name, bytes): - peer_name = peer_name.decode() + peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) + if isinstance(peer_name, bytes): + peer_name = peer_name.decode() - assert ( - peer_name == remote_agent.agent_name - ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" + assert ( + peer_name == remote_agent.agent_name + ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" - page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - kv_page_xfer_handles = self._create_paged_xfer_handles( - page_mem_desc, remote_agent.num_pages, agent_name=peer_name - ) - remote_agent.page_xfer_handles = kv_page_xfer_handles + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + kv_page_xfer_handles = self._create_paged_xfer_handles( + page_mem_desc, remote_agent.num_pages, agent_name=peer_name + ) + remote_agent.page_xfer_handles = kv_page_xfer_handles - logger.info( - f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" - ) + logger.info( + f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" + ) - self.remote_agents[remote_agent.agent_name] = remote_agent + self.remote_agents[remote_agent.agent_name] = remote_agent return def remove_remote_agent(self, peer_name: str): - if peer_name in self.remote_agents: + with self._remote_agents_lock: + remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) + if remote_agent is not None: try: - remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) assert remote_agent.agent_name == peer_name self.nixl_agent.remove_remote_agent(remote_agent.agent_name) if remote_agent.page_xfer_handles is not None: