Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
util.py358 linesDownload Raw Back to mypy
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 
codekingpro/portable-devtools · Team Ai