Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion lightllm/server/api_http_pd.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ async def register_and_keep_alive(websocket: WebSocket):
logger.exception(str(e))
finally:
logger.error(f"client {regist_json} removed")
await g_objs.httpserver_manager.remove_pd(regist_json)
await g_objs.httpserver_manager.remove_pd(regist_json, websocket)
return


Expand Down
53 changes: 47 additions & 6 deletions lightllm/server/httpserver_for_pd_master/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,11 +74,26 @@ def is_healthy(self):
return False

async def register_pd(self, pd_info_json, websocket):
self.pd_manager.register_pd(pd_info_json, websocket)
replaced_node = self.pd_manager.register_pd(pd_info_json, websocket)
if replaced_node is not None:
await self._fail_requests_for_node(replaced_node, "PD node connection was replaced")
return

async def remove_pd(self, pd_info_json):
self.pd_manager.remove_pd(pd_info_json)
async def remove_pd(self, pd_info_json, websocket):
removed_node = self.pd_manager.remove_pd(pd_info_json, websocket)
if removed_node is not None:
await self._fail_requests_for_node(removed_node, "PD node disconnected")
return

async def _fail_requests_for_node(self, node: PD_Client_Obj, reason: str):
error_message = f"{reason}: {node.mode} node {node.client_ip_port}"
affected_requests = [
req_status
for req_status in list(self.req_id_to_out_inf.values())
if req_status.p_node is node or req_status.d_node is node
]
for req_status in affected_requests:
await req_status.fail(RuntimeError(error_message))
return

async def update_req_status(self, upkv_status: PDUpKVStatus):
Expand Down Expand Up @@ -395,6 +410,7 @@ async def fetch_pd_stream(
group_request_id=group_request_id,
stage="prefill",
)
req_status.raise_if_failed()
except ServerBusyError:
logger.warning(f"group_request_id: {group_request_id} wait prefill prompt ids time out")
raise
Expand All @@ -415,6 +431,7 @@ async def fetch_pd_stream(
group_request_id=group_request_id,
stage="decode",
)
req_status.raise_if_failed()
except ServerBusyError:
logger.warning(f"group_request_id: {group_request_id} wait decode stage time out err, server is busy now.")
raise
Expand All @@ -430,6 +447,7 @@ async def fetch_pd_stream(
first_token_gen = False
while True:
await req_status.wait_to_ready()
req_status.raise_if_failed()
if await request.is_disconnected():
raise ClientDisconnected(
group_request_id=group_request_id,
Expand Down Expand Up @@ -636,6 +654,21 @@ def __init__(self, req_id, p_node, d_node) -> None:
self.out_token_info_list: List[Tuple[int, str, dict, FinishStatus]] = []
self.p_node: PD_Client_Obj = p_node
self.d_node: PD_Client_Obj = d_node
self.error: Optional[BaseException] = None

async def fail(self, error: BaseException):
async with self.lock:
if self.error is None:
self.error = error
self.event.set()
self.up_status_event.set()
self.prefill_prompt_ids_event.set()
return

def raise_if_failed(self):
if self.error is not None:
raise self.error
return

async def wait_to_ready(self):
try:
Expand All @@ -646,6 +679,7 @@ async def wait_to_ready(self):
async def can_read(self, req_id_to_out_inf):
async with self.lock:
self.event.clear()
self.raise_if_failed()
assert self.req_id in req_id_to_out_inf, f"error state req_id {self.req_id}"
if len(self.out_token_info_list) == 0:
return False
Expand Down Expand Up @@ -720,6 +754,7 @@ async def check_pd_nodes_health(self):

def register_pd(self, pd_info_json, websocket):
pd_client = PD_Client_Obj(**pd_info_json)
replaced_node = self.url_to_pd_nodes.get(pd_client.client_ip_port)
client_max_req_total_len = pd_client.start_args["max_req_total_len"]
if client_max_req_total_len != self.args.max_req_total_len:
logger.error(
Expand Down Expand Up @@ -756,10 +791,16 @@ def register_pd(self, pd_info_json, websocket):
self.selector.update_nodes(self.prefill_nodes, self.decode_nodes)

logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} registed")
return
if replaced_node is not None and replaced_node.websocket is not websocket:
return replaced_node
return None

def remove_pd(self, pd_info_json):
def remove_pd(self, pd_info_json, websocket):
pd_client = PD_Client_Obj(**pd_info_json)
registered_node = self.url_to_pd_nodes.get(pd_client.client_ip_port)
if registered_node is None or registered_node.websocket is not websocket:
logger.info(f"ignore stale disconnect for {pd_client.mode} node {pd_client.client_ip_port}")
return None

self.url_to_pd_nodes.pop(pd_client.client_ip_port, None)
self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port]
Expand All @@ -768,7 +809,7 @@ def remove_pd(self, pd_info_json):
self.selector.update_nodes(self.prefill_nodes, self.decode_nodes)

logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} removed")
return
return registered_node

def update_node_load_info(self, load_info: Optional[dict]):
"""更新节点负载信息
Expand Down
Loading