From f6258d8caac70429958d9a39ed9e93164f5ef073 Mon Sep 17 00:00:00 2001 From: Viet-Anh Nguyen Date: Sun, 30 Aug 2026 01:58:48 +0700 Subject: [PATCH] fix: clean up workers after model load errors --- .../services/auto_labeling/model_manager.py | 33 ++++++++++--------- anylabeling/utils.py | 10 ++++-- tests/test_model_manager.py | 32 ++++++++++++++++++ tests/test_utils.py | 27 +++++++++++++++ 4 files changed, 85 insertions(+), 17 deletions(-) create mode 100644 tests/test_utils.py diff --git a/anylabeling/services/auto_labeling/model_manager.py b/anylabeling/services/auto_labeling/model_manager.py index ae0a363..769fad7 100644 --- a/anylabeling/services/auto_labeling/model_manager.py +++ b/anylabeling/services/auto_labeling/model_manager.py @@ -417,21 +417,22 @@ def _load_model(self, model_id): model_config = copy.deepcopy(self.model_configs[model_id]) - # Download and extract model - if not model_config.get("has_downloaded", True): - model_config = self._download_and_extract_model(model_config) - if model_config is None: - return - - self.model_configs[model_id].update(model_config) + try: + # Download, extract, and initialize under one error boundary. An + # invalid archive or partially downloaded model must not escape the + # Qt worker slot and strand its thread. + if not model_config.get("has_downloaded", True): + model_config = self._download_and_extract_model(model_config) + if model_config is None: + return - model_type = model_config["type"] - model_class = ModelRegistry.get(model_type) + self.model_configs[model_id].update(model_config) - if not model_class: - raise Exception(f"Unknown model type: {model_type}") + model_type = model_config["type"] + model_class = ModelRegistry.get(model_type) + if not model_class: + raise ValueError(f"Unknown model type: {model_type}") - try: model_config["model"] = model_class( model_config, on_message=self.new_model_status.emit ) @@ -445,9 +446,11 @@ def _load_model(self, model_id): else: self.auto_segmentation_model_unselected.emit() - except Exception as e: # noqa - self.new_model_status.emit(self.tr(f"Error in loading model: {str(e)}")) - print(f"Error in loading model: {str(e)}") + except Exception as error: # noqa + logging.exception("Error loading auto-labeling model") + self.new_model_status.emit( + self.tr("Error in loading model: {error}").format(error=error) + ) return self.loaded_model_config = model_config diff --git a/anylabeling/utils.py b/anylabeling/utils.py index ca14ba0..a49b6af 100644 --- a/anylabeling/utils.py +++ b/anylabeling/utils.py @@ -1,3 +1,5 @@ +import logging + from PyQt6.QtCore import QObject, pyqtSignal, pyqtSlot @@ -12,5 +14,9 @@ def __init__(self, func, *args, **kwargs): @pyqtSlot() def run(self): - self.func(*self.args, **self.kwargs) - self.finished.emit() + try: + self.func(*self.args, **self.kwargs) + except Exception: + logging.exception("Unhandled error in background task") + finally: + self.finished.emit() diff --git a/tests/test_model_manager.py b/tests/test_model_manager.py index cdb42fe..73e6c7a 100644 --- a/tests/test_model_manager.py +++ b/tests/test_model_manager.py @@ -82,6 +82,38 @@ def test_unknown_model_reports_completion(self, mock_load): self.assertEqual(completions, [{}]) + @patch( + "anylabeling.services.auto_labeling.model_manager.ModelManager.load_model_configs" + ) + def test_download_error_is_reported_without_escaping_worker(self, mock_load): + manager = ModelManager() + manager.model_configs = [ + { + "display_name": "Broken download", + "has_downloaded": False, + "type": "yolov8", + } + ] + statuses = [] + completions = [] + manager.new_model_status.connect(statuses.append) + manager.model_loaded.connect(completions.append) + + with ( + patch.object( + manager, + "_download_and_extract_model", + side_effect=ValueError("invalid model archive"), + ), + patch("anylabeling.services.auto_labeling.model_manager.logging.exception"), + ): + result = manager._load_model(0) + + self.assertIsNone(result) + self.assertEqual(statuses, ["Error in loading model: invalid model archive"]) + manager.on_model_download_finished() + self.assertEqual(completions, [{}]) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_utils.py b/tests/test_utils.py new file mode 100644 index 0000000..c2e239d --- /dev/null +++ b/tests/test_utils.py @@ -0,0 +1,27 @@ +"""Tests for background worker cleanup.""" + +import unittest +from unittest import mock + +from anylabeling.utils import GenericWorker + + +class TestGenericWorker(unittest.TestCase): + def test_finished_is_emitted_when_task_raises(self): + finished = [] + + def fail(): + raise RuntimeError("background failure") + + worker = GenericWorker(fail) + worker.finished.connect(lambda: finished.append(True)) + + with mock.patch("anylabeling.utils.logging.exception") as log_exception: + worker.run() + + self.assertEqual(finished, [True]) + log_exception.assert_called_once_with("Unhandled error in background task") + + +if __name__ == "__main__": + unittest.main()