From 5edd7312282018b7646d6916a6b029da18615dbe Mon Sep 17 00:00:00 2001 From: Pushpak Date: Tue, 29 Sep 2026 22:11:07 +0530 Subject: [PATCH] Print missing packages to install when no suitable ImageWriter is found Fixes #7980. resolve_writer raised a generic OptionalImportError when no registered writer backend was available for the given extension, without telling users which optional dependency to install. - OptionalImportError now records the name of the missing package via the new pkg_name argument. - require_pkg passes the missing package name to OptionalImportError. For an installed package with an incompatible version, the recorded name carries the version constraint (e.g. 'torch>=10000'), since an unconstrained 'pip install' would be a no-op in that case. - resolve_writer collects the packages of the unavailable candidate writers and appends installation hints to the error message, e.g. 'No ImageWriter backend found for png. Please install the missing package(s): pillow (e.g. `pip install pillow`).' The PIL import name is translated to the installable pillow distribution; the generic message is kept when dependency details are unavailable. - tests added to tests/data/test_image_rw.py and tests/utils/test_require_pkg.py. Signed-off-by: Pushpak --- monai/data/image_writer.py | 24 +++++++++-- monai/utils/module.py | 23 ++++++++++- tests/data/test_image_rw.py | 71 +++++++++++++++++++++++++++++++++ tests/utils/test_require_pkg.py | 33 +++++++++++++-- 4 files changed, 144 insertions(+), 7 deletions(-) diff --git a/monai/data/image_writer.py b/monai/data/image_writer.py index cc6cdcdead1..d3712f16f6d 100644 --- a/monai/data/image_writer.py +++ b/monai/data/image_writer.py @@ -63,6 +63,9 @@ SUPPORTED_WRITERS: dict = {} +# import names of the writer dependencies that differ from their pip-installable distribution names +_INSTALL_NAMES = {"PIL": "pillow"} + def register_writer(ext_name, *im_writers): """ @@ -99,6 +102,12 @@ def resolve_writer(ext_name, error_if_not_found=True) -> Sequence: As an indexing key it will be converted to a lower case string. error_if_not_found: whether to raise an error if no suitable image writer is found. if True , raise an ``OptionalImportError``, otherwise return an empty tuple. Default is ``True``. + + Raises: + OptionalImportError: When no suitable image writer is found and ``error_if_not_found`` is True. + If the missing writers are known to require packages that are not installed, + the error message additionally suggests the packages to install. + """ if not SUPPORTED_WRITERS: init() @@ -106,19 +115,28 @@ def resolve_writer(ext_name, error_if_not_found=True) -> Sequence: if fmt.startswith("."): fmt = fmt[1:] avail_writers = [] + missing_pkgs = [] default_writers = SUPPORTED_WRITERS.get(EXT_WILDCARD, ()) for _writer in look_up_option(fmt, SUPPORTED_WRITERS, default=default_writers): try: _writer() # this triggers `monai.utils.module.require_pkg` to check the system availability avail_writers.append(_writer) - except OptionalImportError: + except OptionalImportError as e: + if e.pkg_name is not None and e.pkg_name not in missing_pkgs: + missing_pkgs.append(e.pkg_name) continue except Exception: # other writer init errors indicating it exists avail_writers.append(_writer) if not avail_writers and error_if_not_found: - raise OptionalImportError(f"No ImageWriter backend found for {fmt}.") + err_msg = f"No ImageWriter backend found for {fmt}." + if missing_pkgs: + install_names = [_INSTALL_NAMES.get(pkg, pkg) for pkg in missing_pkgs] + install_hints = " or ".join(f"`pip install {name}`" for name in install_names) + err_msg += f" Please install the missing package(s): {' or '.join(install_names)} (e.g. {install_hints})." + raise OptionalImportError(err_msg) writer_tuple = ensure_tuple(avail_writers) - SUPPORTED_WRITERS[fmt] = writer_tuple + if avail_writers: # an empty result is not cached, so a later lookup retries the registered candidates + SUPPORTED_WRITERS[fmt] = writer_tuple return writer_tuple diff --git a/monai/utils/module.py b/monai/utils/module.py index a2569b19fed..fc38545257f 100644 --- a/monai/utils/module.py +++ b/monai/utils/module.py @@ -311,8 +311,18 @@ def __init__(self, required_version, name): class OptionalImportError(ImportError): """ Could not import APIs from an optional dependency. + + Args: + msg: the error message. + pkg_name: name of the missing package that caused the import error, if known. + It is used to provide installation hints to the users. Defaults to ``None``. + """ + def __init__(self, msg: str = "", pkg_name: str | None = None): + super().__init__(msg) + self.pkg_name = pkg_name + def optional_import( module: str, @@ -464,6 +474,12 @@ def require_pkg( raise_error: if True, raise `OptionalImportError` error if the required package is not installed or the version doesn't match requirement, if False, print the error in a warning. + Raises: + OptionalImportError: When ``raise_error`` is True and the required package is not installed or its + version doesn't match the requirement. The error records the name of the package to install + in ``pkg_name``, with the required version constraint appended when ``version`` is specified + (e.g. ``itk>=5.2``), so that installation hints install a compatible version. + """ def _decorator(obj): @@ -476,7 +492,12 @@ def _wrapper(*args, **kwargs): if not has: err_msg = f"required package `{pkg_name}` is not installed or the version doesn't match requirement." if raise_error: - raise OptionalImportError(err_msg) + name = pkg_name + if version: + # record the version constraint so that installation hints install a compatible + # version; `>=` is assumed for any `version_checker` other than `exact_version` + name = f"{pkg_name}{'==' if version_checker is exact_version else '>='}{version}" + raise OptionalImportError(err_msg, pkg_name=name) else: warnings.warn(err_msg) diff --git a/tests/data/test_image_rw.py b/tests/data/test_image_rw.py index d90c1c85711..b2414aa4774 100644 --- a/tests/data/test_image_rw.py +++ b/tests/data/test_image_rw.py @@ -136,6 +136,21 @@ def test_rgb(self, reader, writer): self.png_rw(test_data, reader, writer, np.uint8, False) +class UnavailableWriter: + """ + Simulates a registered writer whose backend dependency is not installed: + ``require_pkg``-decorated writers raise ``OptionalImportError`` on instantiation. + """ + + pkg_name: str | None = "test123" + + def __init__(self): + raise OptionalImportError( + f"required package `{self.pkg_name}` is not installed or the version doesn't match requirement.", + pkg_name=self.pkg_name, + ) + + class TestRegRes(unittest.TestCase): def test_0_default(self): self.assertTrue(len(resolve_writer(".png")) > 0, "has png writer") @@ -150,6 +165,62 @@ def test_1_new(self): register_writer("new2", lambda x: x + 1) self.assertEqual(resolve_writer("new")[0](0), 1) + def test_2_install_hint(self): + register_writer("unknown2", UnavailableWriter) + with self.assertRaises(OptionalImportError) as cm: + resolve_writer("unknown2") + self.assertIn("pip install test123", str(cm.exception)) + + def test_3_install_hint_alias(self): + # the `PIL` import name should be translated to the installable name `pillow` + + class NoPillow(UnavailableWriter): + pkg_name = "PIL" + + register_writer("unknown3", NoPillow) + with self.assertRaises(OptionalImportError) as cm: + resolve_writer("unknown3") + self.assertIn("pip install pillow", str(cm.exception)) + + def test_4_multiple_install_hints(self): + + class NoItk(UnavailableWriter): + pkg_name = "itk" + + class NoNibabel(UnavailableWriter): + pkg_name = "nibabel" + + register_writer("unknown4", NoItk, NoNibabel) + with self.assertRaises(OptionalImportError) as cm: + resolve_writer("unknown4") + self.assertIn("pip install itk", str(cm.exception)) + self.assertIn("pip install nibabel", str(cm.exception)) + + def test_5_no_install_hint(self): + # a writer failing without package details keeps the generic message + + class NoPkgName(UnavailableWriter): + def __init__(self): + raise OptionalImportError("some other reason") + + register_writer("unknown5", NoPkgName) + with self.assertRaises(OptionalImportError) as cm: + resolve_writer("unknown5") + self.assertEqual(str(cm.exception), "No ImageWriter backend found for unknown5.") + + def test_6_empty_result_not_cached(self): + + class NoPkg456(UnavailableWriter): + pkg_name = "test456" + + register_writer("unknown6", NoPkg456) + # a non-raising lookup with no available writer must not cache the empty result, + # so that a later lookup still reports the installation hint + self.assertEqual(resolve_writer("unknown6", error_if_not_found=False), ()) + with self.assertRaises(OptionalImportError) as cm: + resolve_writer("unknown6") + self.assertIn("pip install test456", str(cm.exception)) + @unittest.skipUnless(has_itk, "itk not installed") class TestLoadSaveNrrd(unittest.TestCase): diff --git a/tests/utils/test_require_pkg.py b/tests/utils/test_require_pkg.py index 065a7509a4a..1a5b18a92a5 100644 --- a/tests/utils/test_require_pkg.py +++ b/tests/utils/test_require_pkg.py @@ -13,7 +13,7 @@ import unittest -from monai.utils import OptionalImportError, min_version, require_pkg +from monai.utils import OptionalImportError, exact_version, min_version, require_pkg class TestRequirePkg(unittest.TestCase): @@ -43,7 +43,7 @@ def test_func(x): test_func(x=None) def test_class_exception(self): - with self.assertRaises(OptionalImportError): + with self.assertRaises(OptionalImportError) as cm: @require_pkg(pkg_name="test123") class TestClass: @@ -51,8 +51,11 @@ class TestClass: TestClass() + self.assertEqual(cm.exception.pkg_name, "test123") + def test_class_version_exception(self): - with self.assertRaises(OptionalImportError): + # the installed package is incompatible, so the recorded name carries the version constraint + with self.assertRaises(OptionalImportError) as cm: @require_pkg(pkg_name="torch", version="10000", version_checker=min_version) class TestClass: @@ -60,6 +63,30 @@ class TestClass: TestClass() + self.assertEqual(cm.exception.pkg_name, "torch>=10000") + + def test_func_exact_version_exception(self): + with self.assertRaises(OptionalImportError) as cm: + + @require_pkg(pkg_name="torch", version="10000", version_checker=exact_version) + def test_func(x): + return x + + test_func(x=None) + + self.assertEqual(cm.exception.pkg_name, "torch==10000") + + def test_missing_exact_version_exception(self): + with self.assertRaises(OptionalImportError) as cm: + + @require_pkg(pkg_name="test123", version="1.2", version_checker=exact_version) + def test_func(x): + return x + + test_func(x=None) + + self.assertEqual(cm.exception.pkg_name, "test123==1.2") + def test_func_exception(self): with self.assertRaises(OptionalImportError):