codekingpro/portable-devtools
115k
1"""A mypy_ plugin for managing a number of platform-specific annotations.
2Its functionality can be split into three distinct parts:
3
4* Assigning the (platform-dependent) precisions of certain `~numpy.number`
5 subclasses, including the likes of `~numpy.int_`, `~numpy.intp` and
6 `~numpy.longlong`. See the documentation on
7 :ref:`scalar types <arrays.scalars.built-in>` for a comprehensive overview
8 of the affected classes. Without the plugin the precision of all relevant
9 classes will be inferred as `~typing.Any`.
10* Removing all extended-precision `~numpy.number` subclasses that are
11 unavailable for the platform in question. Most notably this includes the
12 likes of `~numpy.float128` and `~numpy.complex256`. Without the plugin *all*
13 extended-precision types will, as far as mypy is concerned, be available
14 to all platforms.
15* Assigning the (platform-dependent) precision of `~numpy.ctypeslib.c_intp`.
16 Without the plugin the type will default to `ctypes.c_int64`.
17
18 .. versionadded:: 1.22
19
20.. deprecated:: 2.3
21 The :mod:`numpy.typing.mypy_plugin` entry-point is deprecated in favor of
22 platform-agnostic static type inference. Remove
23 ``numpy.typing.mypy_plugin`` from the ``plugins`` section of your mypy
24 configuration; if that surfaces new errors, please open an issue with a
25 minimal reproducer.
26
27Examples
28--------
29To enable the plugin, one must add it to their mypy `configuration file`_:
30
31.. code-block:: ini
32
33 [mypy]
34 plugins = numpy.typing.mypy_plugin
35
36.. _mypy: https://mypy-lang.org/
37.. _configuration file: https://mypy.readthedocs.io/en/stable/config_file.html
38
39"""
40
41from collections.abc import Callable, Iterable
42from typing import TYPE_CHECKING, Final, TypeAlias, cast
43
44import numpy as np
45
46__all__: list[str] = []
47
48
49def _get_precision_dict() -> dict[str, str]:
50 names = [
51 ("_NBitByte", np.byte),
52 ("_NBitShort", np.short),
53 ("_NBitIntC", np.intc),
54 ("_NBitIntP", np.intp),
55 ("_NBitInt", np.int_),
56 ("_NBitLong", np.long),
57 ("_NBitLongLong", np.longlong),
58
59 ("_NBitHalf", np.half),
60 ("_NBitSingle", np.single),
61 ("_NBitDouble", np.double),
62 ("_NBitLongDouble", np.longdouble),
63 ]
64 ret: dict[str, str] = {}
65 for name, typ in names:
66 n = 8 * np.dtype(typ).itemsize
67 ret[f"{_MODULE}._nbit.{name}"] = f"{_MODULE}._nbit_base._{n}Bit"
68 return ret
69
70
71def _get_extended_precision_list() -> list[str]:
72 extended_names = [
73 "float96",
74 "float128",
75 "complex192",
76 "complex256",
77 ]
78 return [i for i in extended_names if hasattr(np, i)]
79
80def _get_c_intp_name() -> str:
81 # Adapted from `np.core._internal._getintp_ctype`
82 return {
83 "i": "c_int",
84 "l": "c_long",
85 "q": "c_longlong",
86 }.get(np.dtype("n").char, "c_long")
87
88
89_MODULE: Final = "numpy._typing"
90
91#: A dictionary mapping type-aliases in `numpy._typing._nbit` to
92#: concrete `numpy.typing.NBitBase` subclasses.
93_PRECISION_DICT: Final = _get_precision_dict()
94
95#: A list with the names of all extended precision `np.number` subclasses.
96_EXTENDED_PRECISION_LIST: Final = _get_extended_precision_list()
97
98#: The name of the ctypes equivalent of `np.intp`
99_C_INTP: Final = _get_c_intp_name()
100
101
102try:
103 if TYPE_CHECKING:
104 from mypy.typeanal import TypeAnalyser
105
106 import mypy.types
107 from mypy.build import PRI_MED
108 from mypy.nodes import ImportFrom, MypyFile, Statement
109 from mypy.plugin import AnalyzeTypeContext, Plugin
110
111except ModuleNotFoundError as e:
112
113 def plugin(version: str) -> type:
114 raise e
115
116else:
117
118 _HookFunc: TypeAlias = Callable[[AnalyzeTypeContext], mypy.types.Type]
119
120 def _hook(ctx: AnalyzeTypeContext) -> mypy.types.Type:
121 """Replace a type-alias with a concrete ``NBitBase`` subclass."""
122 typ, _, api = ctx
123 name = typ.name.split(".")[-1]
124 name_new = _PRECISION_DICT[f"{_MODULE}._nbit.{name}"]
125 return cast("TypeAnalyser", api).named_type(name_new)
126
127 def _index(iterable: Iterable[Statement], id: str) -> int:
128 """Identify the first ``ImportFrom`` instance the specified `id`."""
129 for i, value in enumerate(iterable):
130 if getattr(value, "id", None) == id:
131 return i
132 raise ValueError("Failed to identify a `ImportFrom` instance "
133 f"with the following id: {id!r}")
134
135 def _override_imports(
136 file: MypyFile,
137 module: str,
138 imports: list[tuple[str, str | None]],
139 ) -> None:
140 """Override the first `module`-based import with new `imports`."""
141 # Construct a new `from module import y` statement
142 import_obj = ImportFrom(module, 0, names=imports)
143 import_obj.is_top_level = True
144
145 # Replace the first `module`-based import statement with `import_obj`
146 for lst in [file.defs, cast("list[Statement]", file.imports)]:
147 i = _index(lst, module)
148 lst[i] = import_obj
149
150 class _NumpyPlugin(Plugin):
151 """A mypy plugin for handling versus numpy-specific typing tasks."""
152
153 def get_type_analyze_hook(self, fullname: str) -> _HookFunc | None:
154 """Set the precision of platform-specific `numpy.number`
155 subclasses.
156
157 For example: `numpy.int_`, `numpy.longlong` and `numpy.longdouble`.
158 """
159 if fullname in _PRECISION_DICT:
160 return _hook
161 return None
162
163 def get_additional_deps(
164 self, file: MypyFile
165 ) -> list[tuple[int, str, int]]:
166 """Handle all import-based overrides.
167
168 * Import platform-specific extended-precision `numpy.number`
169 subclasses (*e.g.* `numpy.float96` and `numpy.float128`).
170 * Import the appropriate `ctypes` equivalent to `numpy.intp`.
171
172 """
173 fullname = file.fullname
174 if fullname == "numpy":
175 _override_imports(
176 file,
177 f"{_MODULE}._extended_precision",
178 imports=[(v, v) for v in _EXTENDED_PRECISION_LIST],
179 )
180 elif fullname == "numpy.ctypeslib":
181 _override_imports(
182 file,
183 "ctypes",
184 imports=[(_C_INTP, "_c_intp")],
185 )
186 return [(PRI_MED, fullname, -1)]
187
188 def plugin(version: str) -> type:
189 import warnings
190
191 plugin = "numpy.typing.mypy_plugin"
192 # Deprecated 2025-01-10, NumPy 2.3
193 warn_msg = (
194 f"`{plugin}` is deprecated, and will be removed in a future "
195 f"release. Please remove `plugins = {plugin}` in your mypy config."
196 f"(deprecated in NumPy 2.3)"
197 )
198 warnings.warn(warn_msg, DeprecationWarning, stacklevel=3)
199
200 return _NumpyPlugin
201 