init.detect_modules(): Add base_types filter #80
1 changed files with 9 additions and 4 deletions
|
|
@ -6,16 +6,17 @@ from importlib import import_module
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import MutableMapping
|
from collections.abc import Iterable, MutableMapping
|
||||||
from typing import Any
|
from typing import Any, Sequence
|
||||||
|
|
||||||
def detect_modules(
|
def detect_modules(
|
||||||
namespace: MutableMapping[str, Any],
|
namespace: MutableMapping[str, Any],
|
||||||
prefix: str | None = None,
|
prefix: str | None = None,
|
||||||
skip: set[str] | None = None,
|
skip: set[str] | None = None,
|
||||||
*,
|
*,
|
||||||
|
base_types: Iterable[type[Any]] | None = None,
|
||||||
extend_namespace: bool = True,
|
extend_namespace: bool = True,
|
||||||
) -> list[str]:
|
) -> Sequence[str]:
|
||||||
|
|
||||||
package_name = namespace.get("__name__")
|
package_name = namespace.get("__name__")
|
||||||
package_path = namespace.get("__path__")
|
package_path = namespace.get("__path__")
|
||||||
|
|
@ -41,7 +42,11 @@ def detect_modules(
|
||||||
continue
|
continue
|
||||||
|
|
||||||
module = import_module(f".{module_name}", package_name)
|
module = import_module(f".{module_name}", package_name)
|
||||||
cls = getattr(module, module_name)
|
cls = getattr(module, module_name, None)
|
||||||
|
if cls is None or not isinstance(cls, type):
|
||||||
|
continue
|
||||||
|
if base_types is not None and not issubclass(cls, tuple(base_types)):
|
||||||
|
continue
|
||||||
|
|
||||||
namespace[module_name] = cls
|
namespace[module_name] = cls
|
||||||
ret.append(module_name)
|
ret.append(module_name)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue