_lazyimport.py 5.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176
  1. from __future__ import annotations
  2. __all__ = (
  3. "fix_package_names",
  4. "install_lazy_importer",
  5. "set_deprecated_aliases",
  6. )
  7. import ast
  8. import inspect
  9. import sys
  10. import warnings
  11. from importlib import import_module
  12. from types import ModuleType
  13. from typing import Any
  14. def install_lazy_importer() -> bool:
  15. module_globals = sys._getframe(1).f_globals
  16. module_name = module_globals["__name__"]
  17. module_prefix = module_name + "."
  18. module = sys.modules[module_name]
  19. lazy_map, deprecated_aliases, submodule_names = _build_lazy_map(module)
  20. names = sorted(lazy_map)
  21. # Delete symbols that are not part of the API
  22. del module_globals["TYPE_CHECKING"]
  23. del module_globals["install_lazy_importer"]
  24. if not lazy_map and not deprecated_aliases:
  25. return False
  26. def __getattr__(name: str) -> Any:
  27. if new_name := deprecated_aliases.get(name):
  28. emit_deprecation_warning(module_name, name, new_name)
  29. target_mod, target_attr = new_name.rsplit(".", 1)
  30. elif name in submodule_names:
  31. target_mod, target_attr = "." + name, ""
  32. else:
  33. try:
  34. target_mod, target_attr = lazy_map[name]
  35. except KeyError:
  36. raise AttributeError(
  37. f"module {module_name!r} has no attribute {name!r}"
  38. ) from None
  39. imported = import_module(target_mod, module_name)
  40. value = getattr(imported, target_attr) if target_attr else imported
  41. # patch the module name to match
  42. if (
  43. getattr(value, "__module__", "").startswith(module_prefix)
  44. and name not in deprecated_aliases
  45. ):
  46. value.__module__ = module_name
  47. module_globals[name] = value
  48. return value
  49. def __dir__() -> list[str]:
  50. return names
  51. module_globals["__dir__"] = __dir__
  52. module_globals["__getattr__"] = __getattr__
  53. module_globals.pop("fix_package_names", None)
  54. module_globals.pop("set_deprecated_aliases", None)
  55. return True
  56. def fix_package_names() -> None:
  57. module_globals = sys._getframe(1).f_globals
  58. module_prefix = module_globals["__name__"] + "."
  59. del module_globals[fix_package_names.__name__]
  60. for value in module_globals.values():
  61. if modname := getattr(value, "__module__", ""):
  62. if modname.startswith(module_prefix):
  63. parts = modname.split(".")
  64. value.__module__ = ".".join(
  65. part for part in parts if not part.startswith("_")
  66. )
  67. def emit_deprecation_warning(module_name: str, name: str, target: str) -> None:
  68. warnings.warn(
  69. f"The {module_name}.{name} alias is deprecated, use {target} instead.",
  70. DeprecationWarning,
  71. stacklevel=3,
  72. )
  73. def set_deprecated_aliases(aliases: dict[str, str]) -> None:
  74. module_globals = sys._getframe(1).f_globals
  75. module_name = module_globals["__name__"]
  76. del module_globals[set_deprecated_aliases.__name__]
  77. def __getattr__(name: str) -> Any:
  78. try:
  79. target = aliases[name]
  80. except KeyError:
  81. raise AttributeError(
  82. f"module {module_name!r} has no attribute {name!r}"
  83. ) from None
  84. emit_deprecation_warning(module_name, name, target)
  85. target_modname, attrname = target.rsplit(".", 1)
  86. module = import_module(target_modname)
  87. return getattr(module, attrname)
  88. sys.modules[module_name].__dict__["__getattr__"] = __getattr__
  89. def _build_lazy_map(
  90. module: ModuleType,
  91. ) -> tuple[dict[str, tuple[str, str]], dict[str, str], list[str]]:
  92. try:
  93. source = inspect.getsource(module)
  94. except OSError:
  95. return {}, {}, []
  96. tree = compile(source, module.__file__ or "", "exec", ast.PyCF_ONLY_AST)
  97. assert isinstance(tree, ast.Module)
  98. out: dict[str, tuple[str, str]] = {}
  99. deprecated_aliases: dict[str, str] = {}
  100. submodule_names: list[str] = []
  101. for node in tree.body:
  102. if not isinstance(node, ast.If) or not _is_type_checking_block(node.test):
  103. continue
  104. for stmt in node.body:
  105. match stmt:
  106. case ast.ImportFrom():
  107. if stmt.module is None:
  108. submodule_names.extend(alias.name for alias in stmt.names)
  109. else:
  110. base = "." * stmt.level + (stmt.module or "")
  111. for alias in stmt.names:
  112. if alias.name == "*":
  113. raise RuntimeError("star imports not supported")
  114. exported = alias.asname or alias.name
  115. out[exported] = (base, alias.name)
  116. case ast.Expr() if isinstance(stmt.value, ast.Call):
  117. call = stmt.value
  118. if (
  119. isinstance(call.func, ast.Name)
  120. and call.func.id == "set_deprecated_aliases"
  121. ):
  122. arg0 = call.args[0]
  123. assert isinstance(arg0, ast.Dict)
  124. for key, value in zip(arg0.keys, arg0.values, strict=True):
  125. assert isinstance(key, ast.Constant)
  126. assert isinstance(key.value, str)
  127. assert isinstance(value, ast.Constant)
  128. assert isinstance(value.value, str)
  129. deprecated_aliases[key.value] = value.value
  130. return out, deprecated_aliases, submodule_names
  131. def _is_type_checking_block(test: ast.AST) -> bool:
  132. if not isinstance(test, ast.BoolOp):
  133. return False
  134. subtest = test.values[0]
  135. match subtest:
  136. case ast.Name():
  137. return subtest.id == "TYPE_CHECKING"
  138. case ast.Attribute():
  139. return (
  140. isinstance(subtest.value, ast.Name)
  141. and subtest.value.id == "typing"
  142. and subtest.attr == "TYPE_CHECKING"
  143. )
  144. case _:
  145. return False