diff --git a/bec_widgets/utils/rpc_register.py b/bec_widgets/utils/rpc_register.py index 9e6c0760..f49409f7 100644 --- a/bec_widgets/utils/rpc_register.py +++ b/bec_widgets/utils/rpc_register.py @@ -178,6 +178,19 @@ class RPCRegister: """ self.callbacks.append(callback) + def remove_callback(self, callback: Callable[[dict], None]): + """ + Remove a previously added registry-update callback. Removing a callback + that is not registered is a no-op. + + Args: + callback(Callable[[dict], None]): The callback to be removed. + """ + try: + self.callbacks.remove(callback) + except ValueError: + pass + @classmethod def reset_singleton(cls): """ diff --git a/bec_widgets/utils/rpc_server.py b/bec_widgets/utils/rpc_server.py index 0091a3a8..4f870a23 100644 --- a/bec_widgets/utils/rpc_server.py +++ b/bec_widgets/utils/rpc_server.py @@ -105,7 +105,6 @@ class RPCServer: self._heartbeat_timer = QTimer() self._heartbeat_timer.timeout.connect(self.emit_heartbeat) self._heartbeat_timer.start(200) - self._registry_update_callbacks = [] self._broadcasted_data = {} self._rpc_singleshot_repeats: dict[str, SingleshotRPCRepeat] = {} @@ -529,21 +528,18 @@ class RPCServer: "__rpc__": getattr(connector, "rpc_exposed", True), } - # Suppose clients register callbacks to receive updates - def add_registry_update_callback(self, cb: Callable) -> None: + def shutdown(self): """ - Add a callback to be called whenever the registry is updated. - The specified callback is called whenever the registry is updated. - - Args: - cb (Callable): The callback to be added. It should accept a dictionary of all the - registered RPC objects as an argument. + Shut the RPC server down: stop the heartbeat, release the dispatcher + subscription and the registry callback, and shut the client down. + Safe to call multiple times. """ - self._registry_update_callbacks.append(cb) - - def shutdown(self): # TODO not sure if needed when cleanup is done at level of BECConnector self.status = messages.BECStatus.IDLE self._heartbeat_timer.stop() self.emit_heartbeat() + self.dispatcher.disconnect_slot( + self.on_rpc_update, MessageEndpoints.gui_instructions(self.gui_id) + ) + self.rpc_register.remove_callback(self.broadcast_registry_update) logger.info("Succeeded in shutting down CLI server") self.client.shutdown() diff --git a/tests/unit_tests/test_rpc_server.py b/tests/unit_tests/test_rpc_server.py index fa8ebf14..f21a5325 100644 --- a/tests/unit_tests/test_rpc_server.py +++ b/tests/unit_tests/test_rpc_server.py @@ -297,3 +297,45 @@ def test_run_rpc_delegates_to_rpc_content_class(rpc_server): assert rpc_server.run_rpc(view, "mode", [], {}) == "initial" assert rpc_server.run_rpc(view, "mode", ["creator"], {}) is None assert view.content.mode == "creator" + + +def test_rpc_server_shutdown_releases_registrations(mocked_client): + """Regression test for BW-008/BW-009: shutdown must disconnect the + gui_instructions dispatcher slot and remove the registry callback so a + restarted server does not duplicate registrations and the old server can + be collected.""" + from bec_widgets.utils.rpc_register import RPCRegister + + register = RPCRegister() + callbacks_before = len(register.callbacks) + + server = RPCServer(gui_id="lifecycle_gui", client=mocked_client) + assert len(register.callbacks) == callbacks_before + 1 + + dispatcher_slots_with_topic = [ + slot + for slot in server.dispatcher._registered_slots.values() + if any("lifecycle_gui" in topic for topic in slot.topics) + ] + assert len(dispatcher_slots_with_topic) == 1 + + server.shutdown() + + assert len(register.callbacks) == callbacks_before + dispatcher_slots_with_topic = [ + slot + for slot in server.dispatcher._registered_slots.values() + if any("lifecycle_gui" in topic for topic in slot.topics) + ] + assert dispatcher_slots_with_topic == [] + + # Shutdown must be idempotent. + server.shutdown() + assert len(register.callbacks) == callbacks_before + + +def test_rpc_register_remove_callback_is_noop_for_unknown(rpc_register=None): + from bec_widgets.utils.rpc_register import RPCRegister + + register = RPCRegister() + register.remove_callback(lambda connections: None) # must not raise