codekingpro/portable-devtools
114k
1# util/langhelpers.py
2# Copyright (C) 2005-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
9"""Routines to help with the creation, loading and introspection of
10modules, classes, hierarchies, attributes, functions, and methods.
11
12"""
13from __future__ import annotations
14
15import collections
16import enum
17from functools import update_wrapper
18import inspect
19import itertools
20import operator
21import re
22import sys
23import textwrap
24import threading
25import types
26from types import CodeType
27from typing import Any
28from typing import Callable
29from typing import cast
30from typing import Dict
31from typing import FrozenSet
32from typing import Generic
33from typing import Iterator
34from typing import List
35from typing import Mapping
36from typing import NoReturn
37from typing import Optional
38from typing import overload
39from typing import Sequence
40from typing import Set
41from typing import Tuple
42from typing import Type
43from typing import TYPE_CHECKING
44from typing import TypeVar
45from typing import Union
46import warnings
47
48from . import _collections
49from . import compat
50from ._has_cy import HAS_CYEXTENSION
51from .typing import Literal
52from .. import exc
53
54_T = TypeVar("_T")
55_T_co = TypeVar("_T_co", covariant=True)
56_F = TypeVar("_F", bound=Callable[..., Any])
57_MP = TypeVar("_MP", bound="memoized_property[Any]")
58_MA = TypeVar("_MA", bound="HasMemoized.memoized_attribute[Any]")
59_HP = TypeVar("_HP", bound="hybridproperty[Any]")
60_HM = TypeVar("_HM", bound="hybridmethod[Any]")
61
62
63if compat.py310:
64
65 def get_annotations(obj: Any) -> Mapping[str, Any]:
66 return inspect.get_annotations(obj)
67
68else:
69
70 def get_annotations(obj: Any) -> Mapping[str, Any]:
71 # it's been observed that cls.__annotations__ can be non present.
72 # it's not clear what causes this, running under tox py37/38 it
73 # happens, running straight pytest it doesnt
74
75 # https://docs.python.org/3/howto/annotations.html#annotations-howto
76 if isinstance(obj, type):
77 ann = obj.__dict__.get("__annotations__", None)
78 else:
79 ann = getattr(obj, "__annotations__", None)
80
81 if ann is None:
82 return _collections.EMPTY_DICT
83 else:
84 return cast("Mapping[str, Any]", ann)
85
86
87def md5_hex(x: Any) -> str:
88 x = x.encode("utf-8")
89 m = compat.md5_not_for_security()
90 m.update(x)
91 return cast(str, m.hexdigest())
92
93
94class safe_reraise:
95 """Reraise an exception after invoking some
96 handler code.
97
98 Stores the existing exception info before
99 invoking so that it is maintained across a potential
100 coroutine context switch.
101
102 e.g.::
103
104 try:
105 sess.commit()
106 except:
107 with safe_reraise():
108 sess.rollback()
109
110 TODO: we should at some point evaluate current behaviors in this regard
111 based on current greenlet, gevent/eventlet implementations in Python 3, and
112 also see the degree to which our own asyncio (based on greenlet also) is
113 impacted by this. .rollback() will cause IO / context switch to occur in
114 all these scenarios; what happens to the exception context from an
115 "except:" block if we don't explicitly store it? Original issue was #2703.
116
117 """
118
119 __slots__ = ("_exc_info",)
120
121 _exc_info: Union[
122 None,
123 Tuple[
124 Type[BaseException],
125 BaseException,
126 types.TracebackType,
127 ],
128 Tuple[None, None, None],
129 ]
130
131 def __enter__(self) -> None:
132 self._exc_info = sys.exc_info()
133
134 def __exit__(
135 self,
136 type_: Optional[Type[BaseException]],
137 value: Optional[BaseException],
138 traceback: Optional[types.TracebackType],
139 ) -> NoReturn:
140 assert self._exc_info is not None
141 # see #2703 for notes
142 if type_ is None:
143 exc_type, exc_value, exc_tb = self._exc_info
144 assert exc_value is not None
145 self._exc_info = None # remove potential circular references
146 raise exc_value.with_traceback(exc_tb)
147 else:
148 self._exc_info = None # remove potential circular references
149 assert value is not None
150 raise value.with_traceback(traceback)
151
152
153def walk_subclasses(cls: Type[_T]) -> Iterator[Type[_T]]:
154 seen: Set[Any] = set()
155
156 stack = [cls]
157 while stack:
158 cls = stack.pop()
159 if cls in seen:
160 continue
161 else:
162 seen.add(cls)
163 stack.extend(cls.__subclasses__())
164 yield cls
165
166
167def string_or_unprintable(element: Any) -> str:
168 if isinstance(element, str):
169 return element
170 else:
171 try:
172 return str(element)
173 except Exception:
174 return "unprintable element %r" % element
175
176
177def clsname_as_plain_name(
178 cls: Type[Any], use_name: Optional[str] = None
179) -> str:
180 name = use_name or cls.__name__
181 return " ".join(n.lower() for n in re.findall(r"([A-Z][a-z]+|SQL)", name))
182
183
184def method_is_overridden(
185 instance_or_cls: Union[Type[Any], object],
186 against_method: Callable[..., Any],
187) -> bool:
188 """Return True if the two class methods don't match."""
189
190 if not isinstance(instance_or_cls, type):
191 current_cls = instance_or_cls.__class__
192 else:
193 current_cls = instance_or_cls
194
195 method_name = against_method.__name__
196
197 current_method: types.MethodType = getattr(current_cls, method_name)
198
199 return current_method != against_method
200
201
202def decode_slice(slc: slice) -> Tuple[Any, ...]:
203 """decode a slice object as sent to __getitem__.
204
205 takes into account the 2.5 __index__() method, basically.
206
207 """
208 ret: List[Any] = []
209 for x in slc.start, slc.stop, slc.step:
210 if hasattr(x, "__index__"):
211 x = x.__index__()
212 ret.append(x)
213 return tuple(ret)
214
215
216def _unique_symbols(used: Sequence[str], *bases: str) -> Iterator[str]:
217 used_set = set(used)
218 for base in bases:
219 pool = itertools.chain(
220 (base,),
221 map(lambda i: base + str(i), range(1000)),
222 )
223 for sym in pool:
224 if sym not in used_set:
225 used_set.add(sym)
226 yield sym
227 break
228 else:
229 raise NameError("exhausted namespace for symbol base %s" % base)
230
231
232def map_bits(fn: Callable[[int], Any], n: int) -> Iterator[Any]:
233 """Call the given function given each nonzero bit from n."""
234
235 while n:
236 b = n & (~n + 1)
237 yield fn(b)
238 n ^= b
239
240
241_Fn = TypeVar("_Fn", bound="Callable[..., Any]")
242
243# this seems to be in flux in recent mypy versions
244
245
246def decorator(target: Callable[..., Any]) -> Callable[[_Fn], _Fn]:
247 """A signature-matching decorator factory."""
248
249 def decorate(fn: _Fn) -> _Fn:
250 if not inspect.isfunction(fn) and not inspect.ismethod(fn):
251 raise Exception("not a decoratable function")
252
253 spec = compat.inspect_getfullargspec(fn)
254 env: Dict[str, Any] = {}
255
256 spec = _update_argspec_defaults_into_env(spec, env)
257
258 names = (
259 tuple(cast("Tuple[str, ...]", spec[0]))
260 + cast("Tuple[str, ...]", spec[1:3])
261 + (fn.__name__,)
262 )
263 targ_name, fn_name = _unique_symbols(names, "target", "fn")
264
265 metadata: Dict[str, Optional[str]] = dict(target=targ_name, fn=fn_name)
266 metadata.update(format_argspec_plus(spec, grouped=False))
267 metadata["name"] = fn.__name__
268
269 if inspect.iscoroutinefunction(fn):
270 metadata["prefix"] = "async "
271 metadata["target_prefix"] = "await "
272 else:
273 metadata["prefix"] = ""
274 metadata["target_prefix"] = ""
275
276 # look for __ positional arguments. This is a convention in
277 # SQLAlchemy that arguments should be passed positionally
278 # rather than as keyword
279 # arguments. note that apply_pos doesn't currently work in all cases
280 # such as when a kw-only indicator "*" is present, which is why
281 # we limit the use of this to just that case we can detect. As we add
282 # more kinds of methods that use @decorator, things may have to
283 # be further improved in this area
284 if "__" in repr(spec[0]):
285 code = (
286 """\
287%(prefix)sdef %(name)s%(grouped_args)s:
288 return %(target_prefix)s%(target)s(%(fn)s, %(apply_pos)s)
289"""
290 % metadata
291 )
292 else:
293 code = (
294 """\
295%(prefix)sdef %(name)s%(grouped_args)s:
296 return %(target_prefix)s%(target)s(%(fn)s, %(apply_kw)s)
297"""
298 % metadata
299 )
300
301 mod = sys.modules[fn.__module__]
302 env.update(vars(mod))
303 env.update({targ_name: target, fn_name: fn, "__name__": fn.__module__})
304
305 decorated = cast(
306 types.FunctionType,
307 _exec_code_in_env(code, env, fn.__name__),
308 )
309 decorated.__defaults__ = getattr(fn, "__func__", fn).__defaults__
310
311 decorated.__wrapped__ = fn # type: ignore[attr-defined]
312 return update_wrapper(decorated, fn) # type: ignore[return-value]
313
314 return update_wrapper(decorate, target) # type: ignore[return-value]
315
316
317def _update_argspec_defaults_into_env(spec, env):
318 """given a FullArgSpec, convert defaults to be symbol names in an env."""
319
320 if spec.defaults:
321 new_defaults = []
322 i = 0
323 for arg in spec.defaults:
324 if type(arg).__module__ not in ("builtins", "__builtin__"):
325 name = "x%d" % i
326 env[name] = arg
327 new_defaults.append(name)
328 i += 1
329 else:
330 new_defaults.append(arg)
331 elem = list(spec)
332 elem[3] = tuple(new_defaults)
333 return compat.FullArgSpec(*elem)
334 else:
335 return spec
336
337
338def _exec_code_in_env(
339 code: Union[str, types.CodeType], env: Dict[str, Any], fn_name: str
340) -> Callable[..., Any]:
341 exec(code, env)
342 return env[fn_name] # type: ignore[no-any-return]
343
344
345_PF = TypeVar("_PF")
346_TE = TypeVar("_TE")
347
348
349class PluginLoader:
350 def __init__(
351 self, group: str, auto_fn: Optional[Callable[..., Any]] = None
352 ):
353 self.group = group
354 self.impls: Dict[str, Any] = {}
355 self.auto_fn = auto_fn
356
357 def clear(self):
358 self.impls.clear()
359
360 def load(self, name: str) -> Any:
361 if name in self.impls:
362 return self.impls[name]()
363
364 if self.auto_fn:
365 loader = self.auto_fn(name)
366 if loader:
367 self.impls[name] = loader
368 return loader()
369
370 for impl in compat.importlib_metadata_get(self.group):
371 if impl.name == name:
372 self.impls[name] = impl.load
373 return impl.load()
374
375 raise exc.NoSuchModuleError(
376 "Can't load plugin: %s:%s" % (self.group, name)
377 )
378
379 def register(self, name: str, modulepath: str, objname: str) -> None:
380 def load():
381 mod = __import__(modulepath)
382 for token in modulepath.split(".")[1:]:
383 mod = getattr(mod, token)
384 return getattr(mod, objname)
385
386 self.impls[name] = load
387
388
389def _inspect_func_args(fn):
390 try:
391 co_varkeywords = inspect.CO_VARKEYWORDS
392 except AttributeError:
393 # https://docs.python.org/3/library/inspect.html
394 # The flags are specific to CPython, and may not be defined in other
395 # Python implementations. Furthermore, the flags are an implementation
396 # detail, and can be removed or deprecated in future Python releases.
397 spec = compat.inspect_getfullargspec(fn)
398 return spec[0], bool(spec[2])
399 else:
400 # use fn.__code__ plus flags to reduce method call overhead
401 co = fn.__code__
402 nargs = co.co_argcount
403 return (
404 list(co.co_varnames[:nargs]),
405 bool(co.co_flags & co_varkeywords),
406 )
407
408
409@overload
410def get_cls_kwargs(
411 cls: type,
412 *,
413 _set: Optional[Set[str]] = None,
414 raiseerr: Literal[True] = ...,
415) -> Set[str]: ...
416
417
418@overload
419def get_cls_kwargs(
420 cls: type, *, _set: Optional[Set[str]] = None, raiseerr: bool = False
421) -> Optional[Set[str]]: ...
422
423
424def get_cls_kwargs(
425 cls: type, *, _set: Optional[Set[str]] = None, raiseerr: bool = False
426) -> Optional[Set[str]]:
427 r"""Return the full set of inherited kwargs for the given `cls`.
428
429 Probes a class's __init__ method, collecting all named arguments. If the
430 __init__ defines a \**kwargs catch-all, then the constructor is presumed
431 to pass along unrecognized keywords to its base classes, and the
432 collection process is repeated recursively on each of the bases.
433
434 Uses a subset of inspect.getfullargspec() to cut down on method overhead,
435 as this is used within the Core typing system to create copies of type
436 objects which is a performance-sensitive operation.
437
438 No anonymous tuple arguments please !
439
440 """
441 toplevel = _set is None
442 if toplevel:
443 _set = set()
444 assert _set is not None
445
446 ctr = cls.__dict__.get("__init__", False)
447
448 has_init = (
449 ctr
450 and isinstance(ctr, types.FunctionType)
451 and isinstance(ctr.__code__, types.CodeType)
452 )
453
454 if has_init:
455 names, has_kw = _inspect_func_args(ctr)
456 _set.update(names)
457
458 if not has_kw and not toplevel:
459 if raiseerr:
460 raise TypeError(
461 f"given cls {cls} doesn't have an __init__ method"
462 )
463 else:
464 return None
465 else:
466 has_kw = False
467
468 if not has_init or has_kw:
469 for c in cls.__bases__:
470 if get_cls_kwargs(c, _set=_set) is None:
471 break
472
473 _set.discard("self")
474 return _set
475
476
477def get_func_kwargs(func: Callable[..., Any]) -> List[str]:
478 """Return the set of legal kwargs for the given `func`.
479
480 Uses getargspec so is safe to call for methods, functions,
481 etc.
482
483 """
484
485 return compat.inspect_getfullargspec(func)[0]
486
487
488def get_callable_argspec(
489 fn: Callable[..., Any], no_self: bool = False, _is_init: bool = False
490) -> compat.FullArgSpec:
491 """Return the argument signature for any callable.
492
493 All pure-Python callables are accepted, including
494 functions, methods, classes, objects with __call__;
495 builtins and other edge cases like functools.partial() objects
496 raise a TypeError.
497
498 """
499 if inspect.isbuiltin(fn):
500 raise TypeError("Can't inspect builtin: %s" % fn)
501 elif inspect.isfunction(fn):
502 if _is_init and no_self:
503 spec = compat.inspect_getfullargspec(fn)
504 return compat.FullArgSpec(
505 spec.args[1:],
506 spec.varargs,
507 spec.varkw,
508 spec.defaults,
509 spec.kwonlyargs,
510 spec.kwonlydefaults,
511 spec.annotations,
512 )
513 else:
514 return compat.inspect_getfullargspec(fn)
515 elif inspect.ismethod(fn):
516 if no_self and (_is_init or fn.__self__):
517 spec = compat.inspect_getfullargspec(fn.__func__)
518 return compat.FullArgSpec(
519 spec.args[1:],
520 spec.varargs,
521 spec.varkw,
522 spec.defaults,
523 spec.kwonlyargs,
524 spec.kwonlydefaults,
525 spec.annotations,
526 )
527 else:
528 return compat.inspect_getfullargspec(fn.__func__)
529 elif inspect.isclass(fn):
530 return get_callable_argspec(
531 fn.__init__, no_self=no_self, _is_init=True
532 )
533 elif hasattr(fn, "__func__"):
534 return compat.inspect_getfullargspec(fn.__func__)
535 elif hasattr(fn, "__call__"):
536 if inspect.ismethod(fn.__call__):
537 return get_callable_argspec(fn.__call__, no_self=no_self)
538 else:
539 raise TypeError("Can't inspect callable: %s" % fn)
540 else:
541 raise TypeError("Can't inspect callable: %s" % fn)
542
543
544def format_argspec_plus(
545 fn: Union[Callable[..., Any], compat.FullArgSpec], grouped: bool = True
546) -> Dict[str, Optional[str]]:
547 """Returns a dictionary of formatted, introspected function arguments.
548
549 A enhanced variant of inspect.formatargspec to support code generation.
550
551 fn
552 An inspectable callable or tuple of inspect getargspec() results.
553 grouped
554 Defaults to True; include (parens, around, argument) lists
555
556 Returns:
557
558 args
559 Full inspect.formatargspec for fn
560 self_arg
561 The name of the first positional argument, varargs[0], or None
562 if the function defines no positional arguments.
563 apply_pos
564 args, re-written in calling rather than receiving syntax. Arguments are
565 passed positionally.
566 apply_kw
567 Like apply_pos, except keyword-ish args are passed as keywords.
568 apply_pos_proxied
569 Like apply_pos but omits the self/cls argument
570
571 Example::
572
573 >>> format_argspec_plus(lambda self, a, b, c=3, **d: 123)
574 {'grouped_args': '(self, a, b, c=3, **d)',
575 'self_arg': 'self',
576 'apply_kw': '(self, a, b, c=c, **d)',
577 'apply_pos': '(self, a, b, c, **d)'}
578
579 """
580 if callable(fn):
581 spec = compat.inspect_getfullargspec(fn)
582 else:
583 spec = fn
584
585 args = compat.inspect_formatargspec(*spec)
586
587 apply_pos = compat.inspect_formatargspec(
588 spec[0], spec[1], spec[2], None, spec[4]
589 )
590
591 if spec[0]:
592 self_arg = spec[0][0]
593
594 apply_pos_proxied = compat.inspect_formatargspec(
595 spec[0][1:], spec[1], spec[2], None, spec[4]
596 )
597
598 elif spec[1]:
599 # I'm not sure what this is
600 self_arg = "%s[0]" % spec[1]
601
602 apply_pos_proxied = apply_pos
603 else:
604 self_arg = None
605 apply_pos_proxied = apply_pos
606
607 num_defaults = 0
608 if spec[3]:
609 num_defaults += len(cast(Tuple[Any], spec[3]))
610 if spec[4]:
611 num_defaults += len(spec[4])
612
613 name_args = spec[0] + spec[4]
614
615 defaulted_vals: Union[List[str], Tuple[()]]
616
617 if num_defaults:
618 defaulted_vals = name_args[0 - num_defaults :]
619 else:
620 defaulted_vals = ()
621
622 apply_kw = compat.inspect_formatargspec(
623 name_args,
624 spec[1],
625 spec[2],
626 defaulted_vals,
627 formatvalue=lambda x: "=" + str(x),
628 )
629
630 if spec[0]:
631 apply_kw_proxied = compat.inspect_formatargspec(
632 name_args[1:],
633 spec[1],
634 spec[2],
635 defaulted_vals,
636 formatvalue=lambda x: "=" + str(x),
637 )
638 else:
639 apply_kw_proxied = apply_kw
640
641 if grouped:
642 return dict(
643 grouped_args=args,
644 self_arg=self_arg,
645 apply_pos=apply_pos,
646 apply_kw=apply_kw,
647 apply_pos_proxied=apply_pos_proxied,
648 apply_kw_proxied=apply_kw_proxied,
649 )
650 else:
651 return dict(
652 grouped_args=args,
653 self_arg=self_arg,
654 apply_pos=apply_pos[1:-1],
655 apply_kw=apply_kw[1:-1],
656 apply_pos_proxied=apply_pos_proxied[1:-1],
657 apply_kw_proxied=apply_kw_proxied[1:-1],
658 )
659
660
661def format_argspec_init(method, grouped=True):
662 """format_argspec_plus with considerations for typical __init__ methods
663
664 Wraps format_argspec_plus with error handling strategies for typical
665 __init__ cases::
666
667 object.__init__ -> (self)
668 other unreflectable (usually C) -> (self, *args, **kwargs)
669
670 """
671 if method is object.__init__:
672 grouped_args = "(self)"
673 args = "(self)" if grouped else "self"
674 proxied = "()" if grouped else ""
675 else:
676 try:
677 return format_argspec_plus(method, grouped=grouped)
678 except TypeError:
679 grouped_args = "(self, *args, **kwargs)"
680 args = grouped_args if grouped else "self, *args, **kwargs"
681 proxied = "(*args, **kwargs)" if grouped else "*args, **kwargs"
682 return dict(
683 self_arg="self",
684 grouped_args=grouped_args,
685 apply_pos=args,
686 apply_kw=args,
687 apply_pos_proxied=proxied,
688 apply_kw_proxied=proxied,
689 )
690
691
692def create_proxy_methods(
693 target_cls: Type[Any],
694 target_cls_sphinx_name: str,
695 proxy_cls_sphinx_name: str,
696 classmethods: Sequence[str] = (),
697 methods: Sequence[str] = (),
698 attributes: Sequence[str] = (),
699 use_intermediate_variable: Sequence[str] = (),
700) -> Callable[[_T], _T]:
701 """A class decorator indicating attributes should refer to a proxy
702 class.
703
704 This decorator is now a "marker" that does nothing at runtime. Instead,
705 it is consumed by the tools/generate_proxy_methods.py script to
706 statically generate proxy methods and attributes that are fully
707 recognized by typing tools such as mypy.
708
709 """
710
711 def decorate(cls):
712 return cls
713
714 return decorate
715
716
717def getargspec_init(method):
718 """inspect.getargspec with considerations for typical __init__ methods
719
720 Wraps inspect.getargspec with error handling for typical __init__ cases::
721
722 object.__init__ -> (self)
723 other unreflectable (usually C) -> (self, *args, **kwargs)
724
725 """
726 try:
727 return compat.inspect_getfullargspec(method)
728 except TypeError:
729 if method is object.__init__:
730 return (["self"], None, None, None)
731 else:
732 return (["self"], "args", "kwargs", None)
733
734
735def unbound_method_to_callable(func_or_cls):
736 """Adjust the incoming callable such that a 'self' argument is not
737 required.
738
739 """
740
741 if isinstance(func_or_cls, types.MethodType) and not func_or_cls.__self__:
742 return func_or_cls.__func__
743 else:
744 return func_or_cls
745
746
747def generic_repr(
748 obj: Any,
749 additional_kw: Sequence[Tuple[str, Any]] = (),
750 to_inspect: Optional[Union[object, List[object]]] = None,
751 omit_kwarg: Sequence[str] = (),
752) -> str:
753 """Produce a __repr__() based on direct association of the __init__()
754 specification vs. same-named attributes present.
755
756 """
757 if to_inspect is None:
758 to_inspect = [obj]
759 else:
760 to_inspect = _collections.to_list(to_inspect)
761
762 missing = object()
763
764 pos_args = []
765 kw_args: _collections.OrderedDict[str, Any] = _collections.OrderedDict()
766 vargs = None
767 for i, insp in enumerate(to_inspect):
768 try:
769 spec = compat.inspect_getfullargspec(insp.__init__)
770 except TypeError:
771 continue
772 else:
773 default_len = len(spec.defaults) if spec.defaults else 0
774 if i == 0:
775 if spec.varargs:
776 vargs = spec.varargs
777 if default_len:
778 pos_args.extend(spec.args[1:-default_len])
779 else:
780 pos_args.extend(spec.args[1:])
781 else:
782 kw_args.update(
783 [(arg, missing) for arg in spec.args[1:-default_len]]
784 )
785
786 if default_len:
787 assert spec.defaults
788 kw_args.update(
789 [
790 (arg, default)
791 for arg, default in zip(
792 spec.args[-default_len:], spec.defaults
793 )
794 ]
795 )
796 output: List[str] = []
797
798 output.extend(repr(getattr(obj, arg, None)) for arg in pos_args)
799
800 if vargs is not None and hasattr(obj, vargs):
801 output.extend([repr(val) for val in getattr(obj, vargs)])
802
803 for arg, defval in kw_args.items():
804 if arg in omit_kwarg:
805 continue
806 try:
807 val = getattr(obj, arg, missing)
808 if val is not missing and val != defval:
809 output.append("%s=%r" % (arg, val))
810 except Exception:
811 pass
812
813 if additional_kw:
814 for arg, defval in additional_kw:
815 try:
816 val = getattr(obj, arg, missing)
817 if val is not missing and val != defval:
818 output.append("%s=%r" % (arg, val))
819 except Exception:
820 pass
821
822 return "%s(%s)" % (obj.__class__.__name__, ", ".join(output))
823
824
825class portable_instancemethod:
826 """Turn an instancemethod into a (parent, name) pair
827 to produce a serializable callable.
828
829 """
830
831 __slots__ = "target", "name", "kwargs", "__weakref__"
832
833 def __getstate__(self):
834 return {
835 "target": self.target,
836 "name": self.name,
837 "kwargs": self.kwargs,
838 }
839
840 def __setstate__(self, state):
841 self.target = state["target"]
842 self.name = state["name"]
843 self.kwargs = state.get("kwargs", ())
844
845 def __init__(self, meth, kwargs=()):
846 self.target = meth.__self__
847 self.name = meth.__name__
848 self.kwargs = kwargs
849
850 def __call__(self, *arg, **kw):
851 kw.update(self.kwargs)
852 return getattr(self.target, self.name)(*arg, **kw)
853
854
855def class_hierarchy(cls):
856 """Return an unordered sequence of all classes related to cls.
857
858 Traverses diamond hierarchies.
859
860 Fibs slightly: subclasses of builtin types are not returned. Thus
861 class_hierarchy(class A(object)) returns (A, object), not A plus every
862 class systemwide that derives from object.
863
864 """
865
866 hier = {cls}
867 process = list(cls.__mro__)
868 while process:
869 c = process.pop()
870 bases = (_ for _ in c.__bases__ if _ not in hier)
871
872 for b in bases:
873 process.append(b)
874 hier.add(b)
875
876 if c.__module__ == "builtins" or not hasattr(c, "__subclasses__"):
877 continue
878
879 for s in [
880 _
881 for _ in (
882 c.__subclasses__()
883 if not issubclass(c, type)
884 else c.__subclasses__(c)
885 )
886 if _ not in hier
887 ]:
888 process.append(s)
889 hier.add(s)
890 return list(hier)
891
892
893def iterate_attributes(cls):
894 """iterate all the keys and attributes associated
895 with a class, without using getattr().
896
897 Does not use getattr() so that class-sensitive
898 descriptors (i.e. property.__get__()) are not called.
899
900 """
901 keys = dir(cls)
902 for key in keys:
903 for c in cls.__mro__:
904 if key in c.__dict__:
905 yield (key, c.__dict__[key])
906 break
907
908
909def monkeypatch_proxied_specials(
910 into_cls,
911 from_cls,
912 skip=None,
913 only=None,
914 name="self.proxy",
915 from_instance=None,
916):
917 """Automates delegation of __specials__ for a proxying type."""
918
919 if only:
920 dunders = only
921 else:
922 if skip is None:
923 skip = (
924 "__slots__",
925 "__del__",
926 "__getattribute__",
927 "__metaclass__",
928 "__getstate__",
929 "__setstate__",
930 )
931 dunders = [
932 m
933 for m in dir(from_cls)
934 if (
935 m.startswith("__")
936 and m.endswith("__")
937 and not hasattr(into_cls, m)
938 and m not in skip
939 )
940 ]
941
942 for method in dunders:
943 try:
944 maybe_fn = getattr(from_cls, method)
945 if not hasattr(maybe_fn, "__call__"):
946 continue
947 maybe_fn = getattr(maybe_fn, "__func__", maybe_fn)
948 fn = cast(types.FunctionType, maybe_fn)
949
950 except AttributeError:
951 continue
952 try:
953 spec = compat.inspect_getfullargspec(fn)
954 fn_args = compat.inspect_formatargspec(spec[0])
955 d_args = compat.inspect_formatargspec(spec[0][1:])
956 except TypeError:
957 fn_args = "(self, *args, **kw)"
958 d_args = "(*args, **kw)"
959
960 py = (
961 "def %(method)s%(fn_args)s: "
962 "return %(name)s.%(method)s%(d_args)s" % locals()
963 )
964
965 env: Dict[str, types.FunctionType] = (
966 from_instance is not None and {name: from_instance} or {}
967 )
968 exec(py, env)
969 try:
970 env[method].__defaults__ = fn.__defaults__
971 except AttributeError:
972 pass
973 setattr(into_cls, method, env[method])
974
975
976def methods_equivalent(meth1, meth2):
977 """Return True if the two methods are the same implementation."""
978
979 return getattr(meth1, "__func__", meth1) is getattr(
980 meth2, "__func__", meth2
981 )
982
983
984def as_interface(obj, cls=None, methods=None, required=None):
985 """Ensure basic interface compliance for an instance or dict of callables.
986
987 Checks that ``obj`` implements public methods of ``cls`` or has members
988 listed in ``methods``. If ``required`` is not supplied, implementing at
989 least one interface method is sufficient. Methods present on ``obj`` that
990 are not in the interface are ignored.
991
992 If ``obj`` is a dict and ``dict`` does not meet the interface
993 requirements, the keys of the dictionary are inspected. Keys present in
994 ``obj`` that are not in the interface will raise TypeErrors.
995
996 Raises TypeError if ``obj`` does not meet the interface criteria.
997
998 In all passing cases, an object with callable members is returned. In the
999 simple case, ``obj`` is returned as-is; if dict processing kicks in then
1000 an anonymous class is returned.
1001
1002 obj
1003 A type, instance, or dictionary of callables.
1004 cls
1005 Optional, a type. All public methods of cls are considered the
1006 interface. An ``obj`` instance of cls will always pass, ignoring
1007 ``required``..
1008 methods
1009 Optional, a sequence of method names to consider as the interface.
1010 required
1011 Optional, a sequence of mandatory implementations. If omitted, an
1012 ``obj`` that provides at least one interface method is considered
1013 sufficient. As a convenience, required may be a type, in which case
1014 all public methods of the type are required.
1015
1016 """
1017 if not cls and not methods:
1018 raise TypeError("a class or collection of method names are required")
1019
1020 if isinstance(cls, type) and isinstance(obj, cls):
1021 return obj
1022
1023 interface = set(methods or [m for m in dir(cls) if not m.startswith("_")])
1024 implemented = set(dir(obj))
1025
1026 complies = operator.ge
1027 if isinstance(required, type):
1028 required = interface
1029 elif not required:
1030 required = set()
1031 complies = operator.gt
1032 else:
1033 required = set(required)
1034
1035 if complies(implemented.intersection(interface), required):
1036 return obj
1037
1038 # No dict duck typing here.
1039 if not isinstance(obj, dict):
1040 qualifier = complies is operator.gt and "any of" or "all of"
1041 raise TypeError(
1042 "%r does not implement %s: %s"
1043 % (obj, qualifier, ", ".join(interface))
1044 )
1045
1046 class AnonymousInterface:
1047 """A callable-holding shell."""
1048
1049 if cls:
1050 AnonymousInterface.__name__ = "Anonymous" + cls.__name__
1051 found = set()
1052
1053 for method, impl in dictlike_iteritems(obj):
1054 if method not in interface:
1055 raise TypeError("%r: unknown in this interface" % method)
1056 if not callable(impl):
1057 raise TypeError("%r=%r is not callable" % (method, impl))
1058 setattr(AnonymousInterface, method, staticmethod(impl))
1059 found.add(method)
1060
1061 if complies(found, required):
1062 return AnonymousInterface
1063
1064 raise TypeError(
1065 "dictionary does not contain required keys %s"
1066 % ", ".join(required - found)
1067 )
1068
1069
1070_GFD = TypeVar("_GFD", bound="generic_fn_descriptor[Any]")
1071
1072
1073class generic_fn_descriptor(Generic[_T_co]):
1074 """Descriptor which proxies a function when the attribute is not
1075 present in dict
1076
1077 This superclass is organized in a particular way with "memoized" and
1078 "non-memoized" implementation classes that are hidden from type checkers,
1079 as Mypy seems to not be able to handle seeing multiple kinds of descriptor
1080 classes used for the same attribute.
1081
1082 """
1083
1084 fget: Callable[..., _T_co]
1085 __doc__: Optional[str]
1086 __name__: str
1087
1088 def __init__(self, fget: Callable[..., _T_co], doc: Optional[str] = None):
1089 self.fget = fget
1090 self.__doc__ = doc or fget.__doc__
1091 self.__name__ = fget.__name__
1092
1093 @overload
1094 def __get__(self: _GFD, obj: None, cls: Any) -> _GFD: ...
1095
1096 @overload
1097 def __get__(self, obj: object, cls: Any) -> _T_co: ...
1098
1099 def __get__(self: _GFD, obj: Any, cls: Any) -> Union[_GFD, _T_co]:
1100 raise NotImplementedError()
1101
1102 if TYPE_CHECKING:
1103
1104 def __set__(self, instance: Any, value: Any) -> None: ...
1105
1106 def __delete__(self, instance: Any) -> None: ...
1107
1108 def _reset(self, obj: Any) -> None:
1109 raise NotImplementedError()
1110
1111 @classmethod
1112 def reset(cls, obj: Any, name: str) -> None:
1113 raise NotImplementedError()
1114
1115
1116class _non_memoized_property(generic_fn_descriptor[_T_co]):
1117 """a plain descriptor that proxies a function.
1118
1119 primary rationale is to provide a plain attribute that's
1120 compatible with memoized_property which is also recognized as equivalent
1121 by mypy.
1122
1123 """
1124
1125 if not TYPE_CHECKING:
1126
1127 def __get__(self, obj, cls):
1128 if obj is None:
1129 return self
1130 return self.fget(obj)
1131
1132
1133class _memoized_property(generic_fn_descriptor[_T_co]):
1134 """A read-only @property that is only evaluated once."""
1135
1136 if not TYPE_CHECKING:
1137
1138 def __get__(self, obj, cls):
1139 if obj is None:
1140 return self
1141 obj.__dict__[self.__name__] = result = self.fget(obj)
1142 return result
1143
1144 def _reset(self, obj):
1145 _memoized_property.reset(obj, self.__name__)
1146
1147 @classmethod
1148 def reset(cls, obj, name):
1149 obj.__dict__.pop(name, None)
1150
1151
1152# despite many attempts to get Mypy to recognize an overridden descriptor
1153# where one is memoized and the other isn't, there seems to be no reliable
1154# way other than completely deceiving the type checker into thinking there
1155# is just one single descriptor type everywhere. Otherwise, if a superclass
1156# has non-memoized and subclass has memoized, that requires
1157# "class memoized(non_memoized)". but then if a superclass has memoized and
1158# superclass has non-memoized, the class hierarchy of the descriptors
1159# would need to be reversed; "class non_memoized(memoized)". so there's no
1160# way to achieve this.
1161# additional issues, RO properties:
1162# https://github.com/python/mypy/issues/12440
1163if TYPE_CHECKING:
1164 # allow memoized and non-memoized to be freely mixed by having them
1165 # be the same class
1166 memoized_property = generic_fn_descriptor
1167 non_memoized_property = generic_fn_descriptor
1168
1169 # for read only situations, mypy only sees @property as read only.
1170 # read only is needed when a subtype specializes the return type
1171 # of a property, meaning assignment needs to be disallowed
1172 ro_memoized_property = property
1173 ro_non_memoized_property = property
1174
1175else:
1176 memoized_property = ro_memoized_property = _memoized_property
1177 non_memoized_property = ro_non_memoized_property = _non_memoized_property
1178
1179
1180def memoized_instancemethod(fn: _F) -> _F:
1181 """Decorate a method memoize its return value.
1182
1183 Best applied to no-arg methods: memoization is not sensitive to
1184 argument values, and will always return the same value even when
1185 called with different arguments.
1186
1187 """
1188
1189 def oneshot(self, *args, **kw):
1190 result = fn(self, *args, **kw)
1191
1192 def memo(*a, **kw):
1193 return result
1194
1195 memo.__name__ = fn.__name__
1196 memo.__doc__ = fn.__doc__
1197 self.__dict__[fn.__name__] = memo
1198 return result
1199
1200 return update_wrapper(oneshot, fn) # type: ignore
