codekingpro/portable-devtools
114k
1# ext/mypy/util.py
2# Copyright (C) 2021-2026 the SQLAlchemy authors and contributors
3# <see AUTHORS file>
4#
5# This module is part of SQLAlchemy and is released under
6# the MIT License: https://www.opensource.org/licenses/mit-license.php
7
8from __future__ import annotations
9
10import re
11from typing import Any
12from typing import Iterable
13from typing import Iterator
14from typing import List
15from typing import Optional
16from typing import overload
17from typing import Tuple
18from typing import Type as TypingType
19from typing import TypeVar
20from typing import Union
21
22from mypy import version
23from mypy.messages import format_type as _mypy_format_type
24from mypy.nodes import CallExpr
25from mypy.nodes import ClassDef
26from mypy.nodes import CLASSDEF_NO_INFO
27from mypy.nodes import Context
28from mypy.nodes import Expression
29from mypy.nodes import FuncDef
30from mypy.nodes import IfStmt
31from mypy.nodes import JsonDict
32from mypy.nodes import MemberExpr
33from mypy.nodes import NameExpr
34from mypy.nodes import Statement
35from mypy.nodes import SymbolTableNode
36from mypy.nodes import TypeAlias
37from mypy.nodes import TypeInfo
38from mypy.options import Options
39from mypy.plugin import ClassDefContext
40from mypy.plugin import DynamicClassDefContext
41from mypy.plugin import SemanticAnalyzerPluginInterface
42from mypy.plugins.common import deserialize_and_fixup_type
43from mypy.typeops import map_type_from_supertype
44from mypy.types import CallableType
45from mypy.types import get_proper_type
46from mypy.types import Instance
47from mypy.types import NoneType
48from mypy.types import Type
49from mypy.types import TypeVarType
50from mypy.types import UnboundType
51from mypy.types import UnionType
52
53_vers = tuple(
54 [int(x) for x in version.__version__.split(".") if re.match(r"^\d+$", x)]
55)
56mypy_14 = _vers >= (1, 4)
57
58
59_TArgType = TypeVar("_TArgType", bound=Union[CallExpr, NameExpr])
60
61
62class SQLAlchemyAttribute:
63 def __init__(
64 self,
65 name: str,
66 line: int,
67 column: int,
68 typ: Optional[Type],
69 info: TypeInfo,
70 ) -> None:
71 self.name = name
72 self.line = line
73 self.column = column
74 self.type = typ
75 self.info = info
76
77 def serialize(self) -> JsonDict:
78 assert self.type
79 return {
80 "name": self.name,
81 "line": self.line,
82 "column": self.column,
83 "type": serialize_type(self.type),
84 }
85
86 def expand_typevar_from_subtype(self, sub_type: TypeInfo) -> None:
87 """Expands type vars in the context of a subtype when an attribute is
88 inherited from a generic super type.
89 """
90 if not isinstance(self.type, TypeVarType):
91 return
92
93 self.type = map_type_from_supertype(self.type, sub_type, self.info)
94
95 @classmethod
96 def deserialize(
97 cls,
98 info: TypeInfo,
99 data: JsonDict,
100 api: SemanticAnalyzerPluginInterface,
101 ) -> SQLAlchemyAttribute:
102 data = data.copy()
103 typ = deserialize_and_fixup_type(data.pop("type"), api)
104 return cls(typ=typ, info=info, **data)
105
106
107def name_is_dunder(name: str) -> bool:
108 return bool(re.match(r"^__.+?__$", name))
109
110
111def _set_info_metadata(info: TypeInfo, key: str, data: Any) -> None:
112 info.metadata.setdefault("sqlalchemy", {})[key] = data
113
114
115def _get_info_metadata(info: TypeInfo, key: str) -> Optional[Any]:
116 return info.metadata.get("sqlalchemy", {}).get(key, None)
117
118
119def _get_info_mro_metadata(info: TypeInfo, key: str) -> Optional[Any]:
120 if info.mro:
121 for base in info.mro:
122 metadata = _get_info_metadata(base, key)
123 if metadata is not None:
124 return metadata
125 return None
126
127
128def establish_as_sqlalchemy(info: TypeInfo) -> None:
129 info.metadata.setdefault("sqlalchemy", {})
130
131
132def set_is_base(info: TypeInfo) -> None:
133 _set_info_metadata(info, "is_base", True)
134
135
136def get_is_base(info: TypeInfo) -> bool:
137 is_base = _get_info_metadata(info, "is_base")
138 return is_base is True
139
140
141def has_declarative_base(info: TypeInfo) -> bool:
142 is_base = _get_info_mro_metadata(info, "is_base")
143 return is_base is True
144
145
146def set_has_table(info: TypeInfo) -> None:
147 _set_info_metadata(info, "has_table", True)
148
149
150def get_has_table(info: TypeInfo) -> bool:
151 is_base = _get_info_metadata(info, "has_table")
152 return is_base is True
153
154
155def get_mapped_attributes(
156 info: TypeInfo, api: SemanticAnalyzerPluginInterface
157) -> Optional[List[SQLAlchemyAttribute]]:
158 mapped_attributes: Optional[List[JsonDict]] = _get_info_metadata(
159 info, "mapped_attributes"
160 )
161 if mapped_attributes is None:
162 return None
163
164 attributes: List[SQLAlchemyAttribute] = []
165
166 for data in mapped_attributes:
167 attr = SQLAlchemyAttribute.deserialize(info, data, api)
168 attr.expand_typevar_from_subtype(info)
169 attributes.append(attr)
170
171 return attributes
172
173
174def format_type(typ_: Type, options: Options) -> str:
175 if mypy_14:
176 return _mypy_format_type(typ_, options)
177 else:
178 return _mypy_format_type(typ_) # type: ignore
179
180
181def set_mapped_attributes(
182 info: TypeInfo, attributes: List[SQLAlchemyAttribute]
183) -> None:
184 _set_info_metadata(
185 info,
186 "mapped_attributes",
187 [attribute.serialize() for attribute in attributes],
188 )
189
190
191def fail(api: SemanticAnalyzerPluginInterface, msg: str, ctx: Context) -> None:
192 msg = "[SQLAlchemy Mypy plugin] %s" % msg
193 return api.fail(msg, ctx)
194
195
196def add_global(
197 ctx: Union[ClassDefContext, DynamicClassDefContext],
198 module: str,
199 symbol_name: str,
200 asname: str,
201) -> None:
202 module_globals = ctx.api.modules[ctx.api.cur_mod_id].names
203
204 if asname not in module_globals:
205 lookup_sym: SymbolTableNode = ctx.api.modules[module].names[
206 symbol_name
207 ]
208
209 module_globals[asname] = lookup_sym
210
211
212@overload
213def get_callexpr_kwarg(
214 callexpr: CallExpr, name: str, *, expr_types: None = ...
215) -> Optional[Union[CallExpr, NameExpr]]: ...
216
217
218@overload
219def get_callexpr_kwarg(
220 callexpr: CallExpr,
221 name: str,
222 *,
223 expr_types: Tuple[TypingType[_TArgType], ...],
224) -> Optional[_TArgType]: ...
225
226
227def get_callexpr_kwarg(
228 callexpr: CallExpr,
229 name: str,
230 *,
231 expr_types: Optional[Tuple[TypingType[Any], ...]] = None,
232) -> Optional[Any]:
233 try:
234 arg_idx = callexpr.arg_names.index(name)
235 except ValueError:
236 return None
237
238 kwarg = callexpr.args[arg_idx]
239 if isinstance(
240 kwarg, expr_types if expr_types is not None else (NameExpr, CallExpr)
241 ):
242 return kwarg
243
244 return None
245
246
247def flatten_typechecking(stmts: Iterable[Statement]) -> Iterator[Statement]:
248 for stmt in stmts:
249 if (
250 isinstance(stmt, IfStmt)
251 and isinstance(stmt.expr[0], NameExpr)
252 and stmt.expr[0].fullname == "typing.TYPE_CHECKING"
253 ):
254 yield from stmt.body[0].body
255 else:
256 yield stmt
257
258
259def type_for_callee(callee: Expression) -> Optional[Union[Instance, TypeInfo]]:
260 if isinstance(callee, (MemberExpr, NameExpr)):
261 if isinstance(callee.node, FuncDef):
262 if callee.node.type and isinstance(callee.node.type, CallableType):
263 ret_type = get_proper_type(callee.node.type.ret_type)
264
265 if isinstance(ret_type, Instance):
266 return ret_type
267
268 return None
269 elif isinstance(callee.node, TypeAlias):
270 target_type = get_proper_type(callee.node.target)
271 if isinstance(target_type, Instance):
272 return target_type
273 elif isinstance(callee.node, TypeInfo):
274 return callee.node
275 return None
276
277
278def unbound_to_instance(
279 api: SemanticAnalyzerPluginInterface, typ: Type
280) -> Type:
281 """Take the UnboundType that we seem to get as the ret_type from a FuncDef
282 and convert it into an Instance/TypeInfo kind of structure that seems
283 to work as the left-hand type of an AssignmentStatement.
284
285 """
286
287 if not isinstance(typ, UnboundType):
288 return typ
289
290 # TODO: figure out a more robust way to check this. The node is some
291 # kind of _SpecialForm, there's a typing.Optional that's _SpecialForm,
292 # but I can't figure out how to get them to match up
293 if typ.name == "Optional":
294 # convert from "Optional?" to the more familiar
295 # UnionType[..., NoneType()]
296 return unbound_to_instance(
297 api,
298 UnionType(
299 [unbound_to_instance(api, typ_arg) for typ_arg in typ.args]
300 + [NoneType()]
301 ),
302 )
303
304 node = api.lookup_qualified(typ.name, typ)
305
306 if (
307 node is not None
308 and isinstance(node, SymbolTableNode)
309 and isinstance(node.node, TypeInfo)
310 ):
311 bound_type = node.node
312
313 return Instance(
314 bound_type,
315 [
316 (
317 unbound_to_instance(api, arg)
318 if isinstance(arg, UnboundType)
319 else arg
320 )
321 for arg in typ.args
322 ],
323 )
324 else:
325 return typ
326
327
328def info_for_cls(
329 cls: ClassDef, api: SemanticAnalyzerPluginInterface
330) -> Optional[TypeInfo]:
331 if cls.info is CLASSDEF_NO_INFO:
332 sym = api.lookup_qualified(cls.name, cls)
333 if sym is None:
334 return None
335 assert sym and isinstance(sym.node, TypeInfo)
336 return sym.node
337
338 return cls.info
339
340
341def serialize_type(typ: Type) -> Union[str, JsonDict]:
342 try:
343 return typ.serialize()
344 except Exception:
345 pass
346 if hasattr(typ, "args"):
347 typ.args = tuple(
348 (
349 a.resolve_string_annotation()
350 if hasattr(a, "resolve_string_annotation")
351 else a
352 )
353 for a in typ.args
354 )
355 elif hasattr(typ, "resolve_string_annotation"):
356 typ = typ.resolve_string_annotation()
357 return typ.serialize()
358 