diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 796563325..840063094 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -183,12 +183,6 @@ class CommonKVManager(BaseKVManager): return self.request_status[bootstrap_room] def update_status(self, bootstrap_room: int, status: KVPoll): - if ( - status == KVPoll.Failed - and self.disaggregation_mode == DisaggregationMode.PREFILL - and hasattr(self, "req_to_decode_prefix_len") - ): - self.req_to_decode_prefix_len.pop(bootstrap_room, None) if bootstrap_room not in self.request_status: # Do not resurrect a cleared entry with Failed: once clear() has # popped the room from request_status, any late update_status(Failed) @@ -545,6 +539,13 @@ class CommonKVSender(BaseKVSender): def failure_exception(self): raise Exception("Fake KVReceiver Exception") + def clear(self) -> None: + self.kv_mgr.request_status.pop(self.bootstrap_room, None) + if hasattr(self.kv_mgr, "req_to_decode_prefix_len"): + self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, None) + if hasattr(self.kv_mgr, "transfer_infos"): + self.kv_mgr.transfer_infos.pop(self.bootstrap_room, None) + def abort(self): self.kv_mgr.record_failure( self.bootstrap_room, @@ -740,6 +741,11 @@ class CommonKVReceiver(BaseKVReceiver): def failure_exception(self): raise Exception("Fake KVReceiver Exception") + def clear(self) -> None: + self.kv_mgr.request_status.pop(self.bootstrap_room, None) + self.kv_mgr.required_prefill_response_num_table.pop(self.bootstrap_room, None) + self.kv_mgr.prefill_response_tracker.pop(self.bootstrap_room, None) + def abort(self): self.kv_mgr.record_failure( self.bootstrap_room, diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 5c594027e..7bdd4a601 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -1736,10 +1736,6 @@ class MooncakeKVSender(CommonKVSender): else: return self.conclude_state - def clear(self) -> None: - if self.bootstrap_room in self.kv_mgr.request_status: - self.kv_mgr.request_status.pop(self.bootstrap_room) - def failure_exception(self): # Explicitly set the status to failure since this request has failed in another rank if self.conclude_state is None: @@ -1911,16 +1907,6 @@ class MooncakeKVReceiver(CommonKVReceiver): else: return self.conclude_state - def clear(self) -> None: - if self.bootstrap_room in self.kv_mgr.request_status: - self.kv_mgr.request_status.pop(self.bootstrap_room) - - if self.bootstrap_room in self.kv_mgr.required_prefill_response_num_table: - self.kv_mgr.required_prefill_response_num_table.pop(self.bootstrap_room) - - if self.bootstrap_room in self.kv_mgr.prefill_response_tracker: - self.kv_mgr.prefill_response_tracker.pop(self.bootstrap_room) - def failure_exception(self): # Explicitly set the status to failure since this request has failed in another rank if self.conclude_state is None: diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index c299e8d39..4db7cb0cd 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -959,9 +959,6 @@ class MoriKVSender(CommonKVSender): self._notify_decode(KVPoll.Failed, failure_reason) self.conclude_state = KVPoll.Failed - def clear(self) -> None: - self.kv_mgr.request_status.pop(self.bootstrap_room, None) - def failure_exception(self): if self.conclude_state is None: self._finalize_failure() @@ -1087,9 +1084,7 @@ class MoriKVReceiver(CommonKVReceiver): def clear(self) -> None: if self.bootstrap_room is None: return - self.kv_mgr.request_status.pop(self.bootstrap_room, None) - self.kv_mgr.required_prefill_response_num_table.pop(self.bootstrap_room, None) - self.kv_mgr.prefill_response_tracker.pop(self.bootstrap_room, None) + super().clear() self.kv_mgr._cleanup_room_tracking(self.bootstrap_room) def failure_exception(self): diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index ceca38782..e414116fc 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -1009,9 +1009,6 @@ class NixlKVManager(CommonKVManager): aux_notif, ) handles.append(aux_xfer_handle) - if is_last: - del self.transfer_infos[bootstrap_room] - self.req_to_decode_prefix_len.pop(bootstrap_room, None) return handles def update_transfer_status(self): @@ -1185,7 +1182,6 @@ class NixlKVSender(CommonKVSender): self.chunk_id += 1 if is_last: self.has_sent = True - del self.kv_mgr.request_status[self.bootstrap_room] def poll(self) -> KVPoll: if self._send_failed: