codekingpro/portable-devtools
114k
1# util/typing.py
2# Copyright (C) 2022-2024 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# mypy: allow-untyped-defs, allow-untyped-calls
8
9from __future__ import annotations
10
11import builtins
12import collections.abc as collections_abc
13import re
14import sys
15import typing
16from typing import Any
17from typing import Callable
18from typing import cast
19from typing import Dict
20from typing import ForwardRef
21from typing import Generic
22from typing import Iterable
23from typing import Mapping
24from typing import NewType
25from typing import NoReturn
26from typing import Optional
27from typing import overload
28from typing import Set
29from typing import Tuple
30from typing import Type
31from typing import TYPE_CHECKING
32from typing import TypeVar
33from typing import Union
34
35from . import compat
36
37if True: # zimports removes the tailing comments
38 from typing_extensions import Annotated as Annotated # 3.8
39 from typing_extensions import Concatenate as Concatenate # 3.10
40 from typing_extensions import (
41 dataclass_transform as dataclass_transform, # 3.11,
42 )
43 from typing_extensions import Final as Final # 3.8
44 from typing_extensions import final as final # 3.8
45 from typing_extensions import get_args as get_args # 3.10
46 from typing_extensions import get_origin as get_origin # 3.10
47 from typing_extensions import Literal as Literal # 3.8
48 from typing_extensions import NotRequired as NotRequired # 3.11
49 from typing_extensions import ParamSpec as ParamSpec # 3.10
50 from typing_extensions import Protocol as Protocol # 3.8
51 from typing_extensions import SupportsIndex as SupportsIndex # 3.8
52 from typing_extensions import TypeAlias as TypeAlias # 3.10
53 from typing_extensions import TypedDict as TypedDict # 3.8
54 from typing_extensions import TypeGuard as TypeGuard # 3.10
55 from typing_extensions import Self as Self # 3.11
56 from typing_extensions import TypeAliasType as TypeAliasType # 3.12
57
58_T = TypeVar("_T", bound=Any)
59_KT = TypeVar("_KT")
60_KT_co = TypeVar("_KT_co", covariant=True)
61_KT_contra = TypeVar("_KT_contra", contravariant=True)
62_VT = TypeVar("_VT")
63_VT_co = TypeVar("_VT_co", covariant=True)
64
65
66if compat.py310:
67 # why they took until py310 to put this in stdlib is beyond me,
68 # I've been wanting it since py27
69 from types import NoneType as NoneType
70else:
71 NoneType = type(None) # type: ignore
72
73NoneFwd = ForwardRef("None")
74
75typing_get_args = get_args
76typing_get_origin = get_origin
77
78
79_AnnotationScanType = Union[
80 Type[Any], str, ForwardRef, NewType, TypeAliasType, "GenericProtocol[Any]"
81]
82
83
84class ArgsTypeProcotol(Protocol):
85 """protocol for types that have ``__args__``
86
87 there's no public interface for this AFAIK
88
89 """
90
91 __args__: Tuple[_AnnotationScanType, ...]
92
93
94class GenericProtocol(Protocol[_T]):
95 """protocol for generic types.
96
97 this since Python.typing _GenericAlias is private
98
99 """
100
101 __args__: Tuple[_AnnotationScanType, ...]
102 __origin__: Type[_T]
103
104 # Python's builtin _GenericAlias has this method, however builtins like
105 # list, dict, etc. do not, even though they have ``__origin__`` and
106 # ``__args__``
107 #
108 # def copy_with(self, params: Tuple[_AnnotationScanType, ...]) -> Type[_T]:
109 # ...
110
111
112# copied from TypeShed, required in order to implement
113# MutableMapping.update()
114class SupportsKeysAndGetItem(Protocol[_KT, _VT_co]):
115 def keys(self) -> Iterable[_KT]: ...
116
117 def __getitem__(self, __k: _KT) -> _VT_co: ...
118
119
120# work around https://github.com/microsoft/pyright/issues/3025
121_LiteralStar = Literal["*"]
122
123
124def de_stringify_annotation(
125 cls: Type[Any],
126 annotation: _AnnotationScanType,
127 originating_module: str,
128 locals_: Mapping[str, Any],
129 *,
130 str_cleanup_fn: Optional[Callable[[str, str], str]] = None,
131 include_generic: bool = False,
132 _already_seen: Optional[Set[Any]] = None,
133) -> Type[Any]:
134 """Resolve annotations that may be string based into real objects.
135
136 This is particularly important if a module defines "from __future__ import
137 annotations", as everything inside of __annotations__ is a string. We want
138 to at least have generic containers like ``Mapped``, ``Union``, ``List``,
139 etc.
140
141 """
142 # looked at typing.get_type_hints(), looked at pydantic. We need much
143 # less here, and we here try to not use any private typing internals
144 # or construct ForwardRef objects which is documented as something
145 # that should be avoided.
146
147 original_annotation = annotation
148
149 if is_fwd_ref(annotation):
150 annotation = annotation.__forward_arg__
151
152 if isinstance(annotation, str):
153 if str_cleanup_fn:
154 annotation = str_cleanup_fn(annotation, originating_module)
155
156 annotation = eval_expression(
157 annotation, originating_module, locals_=locals_, in_class=cls
158 )
159
160 if (
161 include_generic
162 and is_generic(annotation)
163 and not is_literal(annotation)
164 ):
165 if _already_seen is None:
166 _already_seen = set()
167
168 if annotation in _already_seen:
169 # only occurs recursively. outermost return type
170 # will always be Type.
171 # the element here will be either ForwardRef or
172 # Optional[ForwardRef]
173 return original_annotation # type: ignore
174 else:
175 _already_seen.add(annotation)
176
177 elements = tuple(
178 de_stringify_annotation(
179 cls,
180 elem,
181 originating_module,
182 locals_,
183 str_cleanup_fn=str_cleanup_fn,
184 include_generic=include_generic,
185 _already_seen=_already_seen,
186 )
187 for elem in annotation.__args__
188 )
189
190 return _copy_generic_annotation_with(annotation, elements)
191 return annotation # type: ignore
192
193
194def _copy_generic_annotation_with(
195 annotation: GenericProtocol[_T], elements: Tuple[_AnnotationScanType, ...]
196) -> Type[_T]:
197 if hasattr(annotation, "copy_with"):
198 # List, Dict, etc. real generics
199 return annotation.copy_with(elements) # type: ignore
200 else:
201 # Python builtins list, dict, etc.
202 return annotation.__origin__[elements] # type: ignore
203
204
205def eval_expression(
206 expression: str,
207 module_name: str,
208 *,
209 locals_: Optional[Mapping[str, Any]] = None,
210 in_class: Optional[Type[Any]] = None,
211) -> Any:
212 try:
213 base_globals: Dict[str, Any] = sys.modules[module_name].__dict__
214 except KeyError as ke:
215 raise NameError(
216 f"Module {module_name} isn't present in sys.modules; can't "
217 f"evaluate expression {expression}"
218 ) from ke
219
220 try:
221 if in_class is not None:
222 cls_namespace = dict(in_class.__dict__)
223 cls_namespace.setdefault(in_class.__name__, in_class)
224
225 # see #10899. We want the locals/globals to take precedence
226 # over the class namespace in this context, even though this
227 # is not the usual way variables would resolve.
228 cls_namespace.update(base_globals)
229
230 annotation = eval(expression, cls_namespace, locals_)
231 else:
232 annotation = eval(expression, base_globals, locals_)
233 except Exception as err:
234 raise NameError(
235 f"Could not de-stringify annotation {expression!r}"
236 ) from err
237 else:
238 return annotation
239
240
241def eval_name_only(
242 name: str,
243 module_name: str,
244 *,
245 locals_: Optional[Mapping[str, Any]] = None,
246) -> Any:
247 if "." in name:
248 return eval_expression(name, module_name, locals_=locals_)
249
250 try:
251 base_globals: Dict[str, Any] = sys.modules[module_name].__dict__
252 except KeyError as ke:
253 raise NameError(
254 f"Module {module_name} isn't present in sys.modules; can't "
255 f"resolve name {name}"
256 ) from ke
257
258 # name only, just look in globals. eval() works perfectly fine here,
259 # however we are seeking to have this be faster, as this occurs for
260 # every Mapper[] keyword, etc. depending on configuration
261 try:
262 return base_globals[name]
263 except KeyError as ke:
264 # check in builtins as well to handle `list`, `set` or `dict`, etc.
265 try:
266 return builtins.__dict__[name]
267 except KeyError:
268 pass
269
270 raise NameError(
271 f"Could not locate name {name} in module {module_name}"
272 ) from ke
273
274
275def resolve_name_to_real_class_name(name: str, module_name: str) -> str:
276 try:
277 obj = eval_name_only(name, module_name)
278 except NameError:
279 return name
280 else:
281 return getattr(obj, "__name__", name)
282
283
284def de_stringify_union_elements(
285 cls: Type[Any],
286 annotation: ArgsTypeProcotol,
287 originating_module: str,
288 locals_: Mapping[str, Any],
289 *,
290 str_cleanup_fn: Optional[Callable[[str, str], str]] = None,
291) -> Type[Any]:
292 return make_union_type(
293 *[
294 de_stringify_annotation(
295 cls,
296 anno,
297 originating_module,
298 {},
299 str_cleanup_fn=str_cleanup_fn,
300 )
301 for anno in annotation.__args__
302 ]
303 )
304
305
306def is_pep593(type_: Optional[_AnnotationScanType]) -> bool:
307 return type_ is not None and typing_get_origin(type_) is Annotated
308
309
310def is_non_string_iterable(obj: Any) -> TypeGuard[Iterable[Any]]:
311 return isinstance(obj, collections_abc.Iterable) and not isinstance(
312 obj, (str, bytes)
313 )
314
315
316def is_literal(type_: _AnnotationScanType) -> bool:
317 return get_origin(type_) is Literal
318
319
320def is_newtype(type_: Optional[_AnnotationScanType]) -> TypeGuard[NewType]:
321 return hasattr(type_, "__supertype__")
322
323 # doesn't work in 3.8, 3.7 as it passes a closure, not an
324 # object instance
325 # return isinstance(type_, NewType)
326
327
328def is_generic(type_: _AnnotationScanType) -> TypeGuard[GenericProtocol[Any]]:
329 return hasattr(type_, "__args__") and hasattr(type_, "__origin__")
330
331
332def is_pep695(type_: _AnnotationScanType) -> TypeGuard[TypeAliasType]:
333 return isinstance(type_, TypeAliasType)
334
335
336def flatten_newtype(type_: NewType) -> Type[Any]:
337 super_type = type_.__supertype__
338 while is_newtype(super_type):
339 super_type = super_type.__supertype__
340 return super_type # type: ignore[return-value]
341
342
343def is_fwd_ref(
344 type_: _AnnotationScanType, check_generic: bool = False
345) -> TypeGuard[ForwardRef]:
346 if isinstance(type_, ForwardRef):
347 return True
348 elif check_generic and is_generic(type_):
349 return any(is_fwd_ref(arg, True) for arg in type_.__args__)
350 else:
351 return False
352
353
354@overload
355def de_optionalize_union_types(type_: str) -> str: ...
356
357
358@overload
359def de_optionalize_union_types(type_: Type[Any]) -> Type[Any]: ...
360
361
362@overload
363def de_optionalize_union_types(
364 type_: _AnnotationScanType,
365) -> _AnnotationScanType: ...
366
367
368def de_optionalize_union_types(
369 type_: _AnnotationScanType,
370) -> _AnnotationScanType:
371 """Given a type, filter out ``Union`` types that include ``NoneType``
372 to not include the ``NoneType``.
373
374 """
375
376 if is_fwd_ref(type_):
377 return de_optionalize_fwd_ref_union_types(type_)
378
379 elif is_optional(type_):
380 typ = set(type_.__args__)
381
382 typ.discard(NoneType)
383 typ.discard(NoneFwd)
384
385 return make_union_type(*typ)
386
387 else:
388 return type_
389
390
391def de_optionalize_fwd_ref_union_types(
392 type_: ForwardRef,
393) -> _AnnotationScanType:
394 """return the non-optional type for Optional[], Union[None, ...], x|None,
395 etc. without de-stringifying forward refs.
396
397 unfortunately this seems to require lots of hardcoded heuristics
398
399 """
400
401 annotation = type_.__forward_arg__
402
403 mm = re.match(r"^(.+?)\[(.+)\]$", annotation)
404 if mm:
405 if mm.group(1) == "Optional":
406 return ForwardRef(mm.group(2))
407 elif mm.group(1) == "Union":
408 elements = re.split(r",\s*", mm.group(2))
409 return make_union_type(
410 *[ForwardRef(elem) for elem in elements if elem != "None"]
411 )
412 else:
413 return type_
414
415 pipe_tokens = re.split(r"\s*\|\s*", annotation)
416 if "None" in pipe_tokens:
417 return ForwardRef("|".join(p for p in pipe_tokens if p != "None"))
418
419 return type_
420
421
422def make_union_type(*types: _AnnotationScanType) -> Type[Any]:
423 """Make a Union type.
424
425 This is needed by :func:`.de_optionalize_union_types` which removes
426 ``NoneType`` from a ``Union``.
427
428 """
429 return cast(Any, Union).__getitem__(types) # type: ignore
430
431
432def expand_unions(
433 type_: Type[Any], include_union: bool = False, discard_none: bool = False
434) -> Tuple[Type[Any], ...]:
435 """Return a type as a tuple of individual types, expanding for
436 ``Union`` types."""
437
438 if is_union(type_):
439 typ = set(type_.__args__)
440
441 if discard_none:
442 typ.discard(NoneType)
443
444 if include_union:
445 return (type_,) + tuple(typ) # type: ignore
446 else:
447 return tuple(typ) # type: ignore
448 else:
449 return (type_,)
450
451
452def is_optional(type_: Any) -> TypeGuard[ArgsTypeProcotol]:
453 return is_origin_of(
454 type_,
455 "Optional",
456 "Union",
457 "UnionType",
458 )
459
460
461def is_optional_union(type_: Any) -> bool:
462 return is_optional(type_) and NoneType in typing_get_args(type_)
463
464
465def is_union(type_: Any) -> TypeGuard[ArgsTypeProcotol]:
466 return is_origin_of(type_, "Union")
467
468
469def is_origin_of_cls(
470 type_: Any, class_obj: Union[Tuple[Type[Any], ...], Type[Any]]
471) -> bool:
472 """return True if the given type has an __origin__ that shares a base
473 with the given class"""
474
475 origin = typing_get_origin(type_)
476 if origin is None:
477 return False
478
479 return isinstance(origin, type) and issubclass(origin, class_obj)
480
481
482def is_origin_of(
483 type_: Any, *names: str, module: Optional[str] = None
484) -> bool:
485 """return True if the given type has an __origin__ with the given name
486 and optional module."""
487
488 origin = typing_get_origin(type_)
489 if origin is None:
490 return False
491
492 return _get_type_name(origin) in names and (
493 module is None or origin.__module__.startswith(module)
494 )
495
496
497def _get_type_name(type_: Type[Any]) -> str:
498 if compat.py310:
499 return type_.__name__
500 else:
501 typ_name = getattr(type_, "__name__", None)
502 if typ_name is None:
503 typ_name = getattr(type_, "_name", None)
504
505 return typ_name # type: ignore
506
507
508class DescriptorProto(Protocol):
509 def __get__(self, instance: object, owner: Any) -> Any: ...
510
511 def __set__(self, instance: Any, value: Any) -> None: ...
512
513 def __delete__(self, instance: Any) -> None: ...
514
515
516_DESC = TypeVar("_DESC", bound=DescriptorProto)
517
518
519class DescriptorReference(Generic[_DESC]):
520 """a descriptor that refers to a descriptor.
521
522 used for cases where we need to have an instance variable referring to an
523 object that is itself a descriptor, which typically confuses typing tools
524 as they don't know when they should use ``__get__`` or not when referring
525 to the descriptor assignment as an instance variable. See
526 sqlalchemy.orm.interfaces.PropComparator.prop
527
528 """
529
530 if TYPE_CHECKING:
531
532 def __get__(self, instance: object, owner: Any) -> _DESC: ...
533
534 def __set__(self, instance: Any, value: _DESC) -> None: ...
535
536 def __delete__(self, instance: Any) -> None: ...
537
538
539_DESC_co = TypeVar("_DESC_co", bound=DescriptorProto, covariant=True)
540
541
542class RODescriptorReference(Generic[_DESC_co]):
543 """a descriptor that refers to a descriptor.
544
545 same as :class:`.DescriptorReference` but is read-only, so that subclasses
546 can define a subtype as the generically contained element
547
548 """
549
550 if TYPE_CHECKING:
551
552 def __get__(self, instance: object, owner: Any) -> _DESC_co: ...
553
554 def __set__(self, instance: Any, value: Any) -> NoReturn: ...
555
556 def __delete__(self, instance: Any) -> NoReturn: ...
557
558
559_FN = TypeVar("_FN", bound=Optional[Callable[..., Any]])
560
561
562class CallableReference(Generic[_FN]):
563 """a descriptor that refers to a callable.
564
565 works around mypy's limitation of not allowing callables assigned
566 as instance variables
567
568
569 """
570
571 if TYPE_CHECKING:
572
573 def __get__(self, instance: object, owner: Any) -> _FN: ...
574
575 def __set__(self, instance: Any, value: _FN) -> None: ...
576
577 def __delete__(self, instance: Any) -> None: ...
578
579
580# $def ro_descriptor_reference(fn: Callable[])
581 