Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/sqlalchemy/util/langhelpers.py: 44%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

1186 statements  

1# util/langhelpers.py 

2# Copyright (C) 2005-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# 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""" 

13 

14from __future__ import annotations 

15 

16import collections 

17import enum 

18from functools import update_wrapper 

19import importlib.metadata 

20import importlib.util 

21import inspect 

22import itertools 

23import linecache 

24import operator 

25import re 

26import sys 

27import textwrap 

28import threading 

29import types 

30from types import CodeType 

31from types import ModuleType 

32from typing import Any 

33from typing import Callable 

34from typing import cast 

35from typing import Dict 

36from typing import FrozenSet 

37from typing import Generic 

38from typing import Iterator 

39from typing import List 

40from typing import Literal 

41from typing import NoReturn 

42from typing import Optional 

43from typing import overload 

44from typing import Sequence 

45from typing import Set 

46from typing import Tuple 

47from typing import Type 

48from typing import TYPE_CHECKING 

49from typing import TypeVar 

50from typing import Union 

51import warnings 

52import weakref 

53 

54from . import _collections 

55from . import compat 

56from .. import exc 

57 

58_T = TypeVar("_T") 

59_T_co = TypeVar("_T_co", covariant=True) 

60_F = TypeVar("_F", bound=Callable[..., Any]) 

61_MA = TypeVar("_MA", bound="HasMemoized.memoized_attribute[Any]") 

62_M = TypeVar("_M", bound=ModuleType) 

63 

64 

65def restore_annotations( 

66 cls: type, new_annotations: dict[str, Any] 

67) -> Callable[[], None]: 

68 """apply alternate annotations to a class, with a callable to restore 

69 the pristine state of the former. 

70 This is used strictly to provide dataclasses on a mapped class, where 

71 in some cases where are making dataclass fields based on an attribute 

72 that is actually a python descriptor on a superclass which we called 

73 to get a value. 

74 if dataclasses were to give us a way to achieve this without swapping 

75 __annotations__, that would be much better. 

76 """ 

77 delattr_ = object() 

78 

79 # pep-649 means classes have "__annotate__", and it's a callable. if it's 

80 # there and is None, we're in "legacy future mode", where it's python 3.14 

81 # or higher and "from __future__ import annotations" is set. in "legacy 

82 # future mode" we have to do the same steps we do for older pythons, 

83 # __annotate__ can be ignored 

84 is_pep649 = hasattr(cls, "__annotate__") and cls.__annotate__ is not None 

85 

86 if is_pep649: 

87 memoized = { 

88 "__annotate__": getattr(cls, "__annotate__", delattr_), 

89 } 

90 else: 

91 memoized = { 

92 "__annotations__": getattr(cls, "__annotations__", delattr_) 

93 } 

94 

95 cls.__annotations__ = new_annotations 

96 

97 def restore(): 

98 for k, v in memoized.items(): 

99 if v is delattr_: 

100 delattr(cls, k) 

101 else: 

102 setattr(cls, k, v) 

103 

104 return restore 

105 

106 

107def md5_hex(x: Any) -> str: 

108 x = x.encode("utf-8") 

109 m = compat.md5_not_for_security() 

110 m.update(x) 

111 return cast(str, m.hexdigest()) 

112 

113 

114class safe_reraise: 

115 """Reraise an exception after invoking some 

116 handler code. 

117 

118 Stores the existing exception info before 

119 invoking so that it is maintained across a potential 

120 coroutine context switch. 

121 

122 e.g.:: 

123 

124 try: 

125 sess.commit() 

126 except: 

127 with safe_reraise(): 

128 sess.rollback() 

129 

130 TODO: we should at some point evaluate current behaviors in this regard 

131 based on current greenlet, gevent/eventlet implementations in Python 3, and 

132 also see the degree to which our own asyncio (based on greenlet also) is 

133 impacted by this. .rollback() will cause IO / context switch to occur in 

134 all these scenarios; what happens to the exception context from an 

135 "except:" block if we don't explicitly store it? Original issue was #2703. 

136 

137 """ 

138 

139 __slots__ = ("_exc_info",) 

140 

141 _exc_info: Union[ 

142 None, 

143 Tuple[ 

144 Type[BaseException], 

145 BaseException, 

146 types.TracebackType, 

147 ], 

148 Tuple[None, None, None], 

149 ] 

150 

151 def __enter__(self) -> None: 

152 self._exc_info = sys.exc_info() 

153 

154 def __exit__( 

155 self, 

156 type_: Optional[Type[BaseException]], 

157 value: Optional[BaseException], 

158 traceback: Optional[types.TracebackType], 

159 ) -> NoReturn: 

160 assert self._exc_info is not None 

161 # see #2703 for notes 

162 if type_ is None: 

163 exc_type, exc_value, exc_tb = self._exc_info 

164 assert exc_value is not None 

165 self._exc_info = None # remove potential circular references 

166 raise exc_value.with_traceback(exc_tb) 

167 else: 

168 self._exc_info = None # remove potential circular references 

169 assert value is not None 

170 raise value.with_traceback(traceback) 

171 

172 

173def walk_subclasses(cls: Type[_T]) -> Iterator[Type[_T]]: 

174 seen: Set[Any] = set() 

175 

176 stack = [cls] 

177 while stack: 

178 cls = stack.pop() 

179 if cls in seen: 

180 continue 

181 else: 

182 seen.add(cls) 

183 stack.extend(cls.__subclasses__()) 

184 yield cls 

185 

186 

187def string_or_unprintable(element: Any) -> str: 

188 if isinstance(element, str): 

189 return element 

190 else: 

191 try: 

192 return str(element) 

193 except Exception: 

194 return "unprintable element %r" % element 

195 

196 

197def clsname_as_plain_name( 

198 cls: Type[Any], use_name: Optional[str] = None 

199) -> str: 

200 name = use_name or cls.__name__ 

201 return " ".join(n.lower() for n in re.findall(r"([A-Z][a-z]+|SQL)", name)) 

202 

203 

204def method_is_overridden( 

205 instance_or_cls: Union[Type[Any], object], 

206 against_method: Callable[..., Any], 

207) -> bool: 

208 """Return True if the two class methods don't match.""" 

209 

210 if not isinstance(instance_or_cls, type): 

211 current_cls = instance_or_cls.__class__ 

212 else: 

213 current_cls = instance_or_cls 

214 

215 method_name = against_method.__name__ 

216 

217 current_method: types.MethodType = getattr(current_cls, method_name) 

218 

219 return current_method != against_method 

220 

221 

222def decode_slice(slc: slice) -> Tuple[Any, ...]: 

223 """decode a slice object as sent to __getitem__. 

224 

225 takes into account the 2.5 __index__() method, basically. 

226 

227 """ 

228 ret: List[Any] = [] 

229 for x in slc.start, slc.stop, slc.step: 

230 if hasattr(x, "__index__"): 

231 x = x.__index__() 

232 ret.append(x) 

233 return tuple(ret) 

234 

235 

236def _unique_symbols(used: Sequence[str], *bases: str) -> Iterator[str]: 

237 used_set = set(used) 

238 for base in bases: 

239 pool = itertools.chain( 

240 (base,), 

241 map(lambda i: base + str(i), range(1000)), 

242 ) 

243 for sym in pool: 

244 if sym not in used_set: 

245 used_set.add(sym) 

246 yield sym 

247 break 

248 else: 

249 raise NameError("exhausted namespace for symbol base %s" % base) 

250 

251 

252def map_bits(fn: Callable[[int], Any], n: int) -> Iterator[Any]: 

253 """Call the given function given each nonzero bit from n.""" 

254 

255 while n: 

256 b = n & (~n + 1) 

257 yield fn(b) 

258 n ^= b 

259 

260 

261_Fn = TypeVar("_Fn", bound="Callable[..., Any]") 

262 

263# this seems to be in flux in recent mypy versions 

264 

265 

266def decorator(target: Callable[..., Any]) -> Callable[[_Fn], _Fn]: 

267 """A signature-matching decorator factory.""" 

268 

269 def decorate(fn: _Fn) -> _Fn: 

270 if not inspect.isfunction(fn) and not inspect.ismethod(fn): 

271 raise Exception("not a decoratable function") 

272 

273 # Python 3.14 defer creating __annotations__ until its used. 

274 # We do not want to create __annotations__ now. 

275 annofunc = getattr(fn, "__annotate__", None) 

276 if annofunc is not None: 

277 fn.__annotate__ = None # type: ignore[union-attr] 

278 try: 

279 spec = compat.inspect_getfullargspec(fn) 

280 finally: 

281 fn.__annotate__ = annofunc # type: ignore[union-attr] 

282 else: 

283 spec = compat.inspect_getfullargspec(fn) 

284 

285 # Do not generate code for annotations. 

286 # update_wrapper() copies the annotation from fn to decorated. 

287 # We use dummy defaults for code generation to avoid having 

288 # copy of large globals for compiling. 

289 # We copy __defaults__ and __kwdefaults__ from fn to decorated. 

290 empty_defaults = (None,) * len(spec.defaults or ()) 

291 empty_kwdefaults = dict.fromkeys(spec.kwonlydefaults or ()) 

292 spec = spec._replace( 

293 annotations={}, 

294 defaults=empty_defaults, 

295 kwonlydefaults=empty_kwdefaults, 

296 ) 

297 

298 names = ( 

299 tuple(cast("Tuple[str, ...]", spec[0])) 

300 + cast("Tuple[str, ...]", spec[1:3]) 

301 + (fn.__name__,) 

302 ) 

303 targ_name, fn_name = _unique_symbols(names, "target", "fn") 

304 

305 metadata: Dict[str, Optional[str]] = dict(target=targ_name, fn=fn_name) 

306 metadata.update(format_argspec_plus(spec, grouped=False)) 

307 metadata["name"] = fn.__name__ 

308 

309 if inspect.iscoroutinefunction(fn): 

310 metadata["prefix"] = "async " 

311 metadata["target_prefix"] = "await " 

312 metadata["target_suffix"] = "" 

313 elif inspect.isgeneratorfunction(fn): 

314 # a generator function has to remain a generator function 

315 # after decoration; tools such as pytest fixtures test for 

316 # inspect.isgeneratorfunction() and will otherwise never 

317 # iterate the function at all 

318 metadata["prefix"] = "" 

319 metadata["target_prefix"] = "(yield from " 

320 metadata["target_suffix"] = ")" 

321 else: 

322 metadata["prefix"] = "" 

323 metadata["target_prefix"] = "" 

324 metadata["target_suffix"] = "" 

325 

326 # look for __ positional arguments. This is a convention in 

327 # SQLAlchemy that arguments should be passed positionally 

328 # rather than as keyword 

329 # arguments. note that apply_pos doesn't currently work in all cases 

330 # such as when a kw-only indicator "*" is present, which is why 

331 # we limit the use of this to just that case we can detect. As we add 

332 # more kinds of methods that use @decorator, things may have to 

333 # be further improved in this area 

334 if "__" in repr(spec[0]): 

335 code = """\ 

336%(prefix)sdef %(name)s%(grouped_args)s: 

337 return %(target_prefix)s%(target)s(%(fn)s, %(apply_pos)s)%(target_suffix)s 

338""" % metadata 

339 else: 

340 code = """\ 

341%(prefix)sdef %(name)s%(grouped_args)s: 

342 return %(target_prefix)s%(target)s(%(fn)s, %(apply_kw)s)%(target_suffix)s 

343""" % metadata 

344 

345 env: Dict[str, Any] = { 

346 targ_name: target, 

347 fn_name: fn, 

348 "__name__": fn.__module__, 

349 } 

350 

351 # the target's name is part of the description because decorators 

352 # built here do get stacked (see ValuesBase.values()), and 

353 # update_wrapper() gives every layer the same __qualname__; without 

354 # it the outer layer would claim the inner layer's source. 

355 decorated = cast( 

356 types.FunctionType, 

357 exec_code_in_env( 

358 code, 

359 env, 

360 fn.__name__, 

361 f"{target.__name__}() wrapper for " 

362 f"{fn.__module__}.{fn.__qualname__}", 

363 ), 

364 ) 

365 decorated.__defaults__ = fn.__defaults__ 

366 decorated.__kwdefaults__ = fn.__kwdefaults__ # type: ignore[union-attr] # noqa: E501 

367 return update_wrapper(decorated, fn) # type: ignore[return-value] 

368 

369 return update_wrapper(decorate, target) # type: ignore[return-value] 

370 

371 

372_LinecacheEntry = Tuple[int, None, List[str], str] 

373 

374 

375def _linecache_cache_getter(): 

376 """safe getter for linecache.cache 

377 

378 linecache.cache despite being non-underscored and widely used is 

379 nonetheless not documented by cPython. Therefore we cannot trust that it 

380 it's present in third party Python distributions, or that it wont 

381 suddenly be removed or changed. Guard against this such that linecache 

382 features will be silently disabled if this should happen. Unit tests in 

383 test/base/test_utils.py ensures linecache.cache remains available for new 

384 releases. 

385 

386 """ 

387 try: 

388 linecache_cache = linecache.cache 

389 except AttributeError: 

390 raise 

391 else: 

392 if not isinstance(linecache_cache, dict): 

393 raise AttributeError("linecache has changed from being a dict") 

394 return linecache_cache 

395 

396 

397def _remove_linecache_entry(filename: str, entry: _LinecacheEntry) -> None: 

398 """Discard a ``linecache`` entry made by :func:`.exec_code_in_env`, if 

399 it is still the live one. 

400 

401 Runs from a weakref finalizer. 

402 

403 This function is only established if we actually added an entry to the 

404 linecache within exec_code_in_env. 

405 

406 """ 

407 try: 

408 linecache_cache = _linecache_cache_getter() 

409 except AttributeError: 

410 # linecache.cache is gone or is no longer a dict; as this runs 

411 # from a finalizer, an exception here would only surface as an 

412 # unraisable error 

413 return 

414 

415 if linecache_cache.get(filename) is entry: 

416 linecache_cache.pop(filename, None) 

417 

418 

419def exec_code_in_env( 

420 code: Union[str, types.CodeType], 

421 env: Dict[str, Any], 

422 fn_name: str, 

423 description: Optional[str] = None, 

424) -> Callable[..., Any]: 

425 """Exec generated ``code`` in ``env`` and return the function it defines. 

426 

427 If ``description`` is passed, the code is compiled against a synthetic 

428 filename which is registered with :mod:`linecache`, allowing traceback 

429 frames to reference the actual source code being referenced. 

430 

431 entries are placed in the cache without an expiration time and a 

432 weakref.finalize() is applied to the function to remove the linecache 

433 entry if and when the function is garbage collected. 

434 

435 """ 

436 filename: Optional[str] = None 

437 entry: Optional[_LinecacheEntry] = None 

438 

439 if description is not None: 

440 assert isinstance(code, str), ( 

441 "a description is only meaningful for source that has not " 

442 "already been compiled" 

443 ) 

444 filename = f"<sqlalchemy generated {description}>" 

445 

446 try: 

447 linecache_cache = _linecache_cache_getter() 

448 except AttributeError: 

449 pass 

450 else: 

451 entry = (len(code), None, code.splitlines(True), filename) 

452 linecache_cache[filename] = entry 

453 code = compile(code, filename, "exec") 

454 

455 exec(code, env) 

456 fn = env[fn_name] 

457 

458 if filename is not None and entry is not None: 

459 # apply a finalizer that addresses on-the-fly functions, ORM mapped 

460 # classes, etc. which might be garbage collected 

461 finalizer = weakref.finalize( 

462 fn, _remove_linecache_entry, filename, entry 

463 ) 

464 # atexit is a property on the C implementation; typeshed 

465 # renders finalize with an empty __slots__ 

466 finalizer.atexit = False # type: ignore[misc] 

467 

468 return fn # type: ignore[no-any-return] 

469 

470 

471_PF = TypeVar("_PF") 

472_TE = TypeVar("_TE") 

473 

474 

475class PluginLoader: 

476 def __init__( 

477 self, group: str, auto_fn: Optional[Callable[..., Any]] = None 

478 ): 

479 self.group = group 

480 self.impls: Dict[str, Any] = {} 

481 self.auto_fn = auto_fn 

482 

483 def clear(self): 

484 self.impls.clear() 

485 

486 def load(self, name: str) -> Any: 

487 if name in self.impls: 

488 return self.impls[name]() 

489 

490 if self.auto_fn: 

491 loader = self.auto_fn(name) 

492 if loader: 

493 self.impls[name] = loader 

494 return loader() 

495 

496 for impl in compat.importlib_metadata_get(self.group): 

497 if impl.name == name: 

498 self.impls[name] = impl.load 

499 return impl.load() 

500 

501 raise exc.NoSuchModuleError( 

502 "Can't load plugin: %s:%s" % (self.group, name) 

503 ) 

504 

505 def register(self, name: str, modulepath: str, objname: str) -> None: 

506 def load(): 

507 mod = __import__(modulepath) 

508 for token in modulepath.split(".")[1:]: 

509 mod = getattr(mod, token) 

510 return getattr(mod, objname) 

511 

512 self.impls[name] = load 

513 

514 def deregister(self, name: str) -> None: 

515 del self.impls[name] 

516 

517 

518def _inspect_func_args(fn): 

519 try: 

520 co_varkeywords = inspect.CO_VARKEYWORDS 

521 except AttributeError: 

522 # https://docs.python.org/3/library/inspect.html 

523 # The flags are specific to CPython, and may not be defined in other 

524 # Python implementations. Furthermore, the flags are an implementation 

525 # detail, and can be removed or deprecated in future Python releases. 

526 spec = compat.inspect_getfullargspec(fn) 

527 return spec[0], bool(spec[2]) 

528 else: 

529 # use fn.__code__ plus flags to reduce method call overhead 

530 co = fn.__code__ 

531 nargs = co.co_argcount 

532 return ( 

533 list(co.co_varnames[:nargs]), 

534 bool(co.co_flags & co_varkeywords), 

535 ) 

536 

537 

538@overload 

539def get_cls_kwargs( 

540 cls: type, 

541 *, 

542 _set: Optional[Set[str]] = None, 

543 raiseerr: Literal[True] = ..., 

544) -> Set[str]: ... 

545 

546 

547@overload 

548def get_cls_kwargs( 

549 cls: type, *, _set: Optional[Set[str]] = None, raiseerr: bool = False 

550) -> Optional[Set[str]]: ... 

551 

552 

553def get_cls_kwargs( 

554 cls: type, *, _set: Optional[Set[str]] = None, raiseerr: bool = False 

555) -> Optional[Set[str]]: 

556 r"""Return the full set of inherited kwargs for the given `cls`. 

557 

558 Probes a class's __init__ method, collecting all named arguments. If the 

559 __init__ defines a \**kwargs catch-all, then the constructor is presumed 

560 to pass along unrecognized keywords to its base classes, and the 

561 collection process is repeated recursively on each of the bases. 

562 

563 Uses a subset of inspect.getfullargspec() to cut down on method overhead, 

564 as this is used within the Core typing system to create copies of type 

565 objects which is a performance-sensitive operation. 

566 

567 No anonymous tuple arguments please ! 

568 

569 """ 

570 toplevel = _set is None 

571 if toplevel: 

572 _set = set() 

573 assert _set is not None 

574 

575 ctr = cls.__dict__.get("__init__", False) 

576 

577 has_init = ( 

578 ctr 

579 and isinstance(ctr, types.FunctionType) 

580 and isinstance(ctr.__code__, types.CodeType) 

581 ) 

582 

583 if has_init: 

584 names, has_kw = _inspect_func_args(ctr) 

585 _set.update(names) 

586 

587 if not has_kw and not toplevel: 

588 if raiseerr: 

589 raise TypeError( 

590 f"given cls {cls} doesn't have an __init__ method" 

591 ) 

592 else: 

593 return None 

594 else: 

595 has_kw = False 

596 

597 if not has_init or has_kw: 

598 for c in cls.__bases__: 

599 if get_cls_kwargs(c, _set=_set) is None: 

600 break 

601 

602 _set.discard("self") 

603 return _set 

604 

605 

606def get_func_kwargs(func: Callable[..., Any]) -> List[str]: 

607 """Return the set of legal kwargs for the given `func`. 

608 

609 Uses getargspec so is safe to call for methods, functions, 

610 etc. 

611 

612 """ 

613 

614 return compat.inspect_getfullargspec(func)[0] 

615 

616 

617def get_callable_argspec( 

618 fn: Callable[..., Any], no_self: bool = False, _is_init: bool = False 

619) -> compat.FullArgSpec: 

620 """Return the argument signature for any callable. 

621 

622 All pure-Python callables are accepted, including 

623 functions, methods, classes, objects with __call__; 

624 builtins and other edge cases like functools.partial() objects 

625 raise a TypeError. 

626 

627 """ 

628 if inspect.isbuiltin(fn): 

629 raise TypeError("Can't inspect builtin: %s" % fn) 

630 elif inspect.isfunction(fn) or ( 

631 hasattr(fn, "__code__") 

632 and not inspect.isclass(fn) 

633 and not inspect.ismethod(fn) 

634 ): 

635 if _is_init and no_self: 

636 spec = compat.inspect_getfullargspec(fn) 

637 return compat.FullArgSpec( 

638 spec.args[1:], 

639 spec.varargs, 

640 spec.varkw, 

641 spec.defaults, 

642 spec.kwonlyargs, 

643 spec.kwonlydefaults, 

644 spec.annotations, 

645 ) 

646 else: 

647 return compat.inspect_getfullargspec(fn) 

648 elif inspect.ismethod(fn): 

649 if no_self and (_is_init or fn.__self__): 

650 spec = compat.inspect_getfullargspec(fn.__func__) 

651 return compat.FullArgSpec( 

652 spec.args[1:], 

653 spec.varargs, 

654 spec.varkw, 

655 spec.defaults, 

656 spec.kwonlyargs, 

657 spec.kwonlydefaults, 

658 spec.annotations, 

659 ) 

660 else: 

661 return compat.inspect_getfullargspec(fn.__func__) 

662 elif inspect.isclass(fn): 

663 return get_callable_argspec( 

664 fn.__init__, no_self=no_self, _is_init=True 

665 ) 

666 elif hasattr(fn, "__func__"): 

667 return compat.inspect_getfullargspec(fn.__func__) 

668 elif hasattr(fn, "__call__"): 

669 if inspect.ismethod(fn.__call__): 

670 return get_callable_argspec(fn.__call__, no_self=no_self) 

671 else: 

672 raise TypeError("Can't inspect callable: %s" % fn) 

673 else: 

674 raise TypeError("Can't inspect callable: %s" % fn) 

675 

676 

677def format_argspec_plus( 

678 fn: Union[Callable[..., Any], compat.FullArgSpec], grouped: bool = True 

679) -> Dict[str, Optional[str]]: 

680 """Returns a dictionary of formatted, introspected function arguments. 

681 

682 A enhanced variant of inspect.formatargspec to support code generation. 

683 

684 fn 

685 An inspectable callable or tuple of inspect getargspec() results. 

686 grouped 

687 Defaults to True; include (parens, around, argument) lists 

688 

689 Returns: 

690 

691 args 

692 Full inspect.formatargspec for fn 

693 self_arg 

694 The name of the first positional argument, varargs[0], or None 

695 if the function defines no positional arguments. 

696 apply_pos 

697 args, re-written in calling rather than receiving syntax. Arguments are 

698 passed positionally. 

699 apply_kw 

700 Like apply_pos, except keyword-ish args are passed as keywords. 

701 apply_pos_proxied 

702 Like apply_pos but omits the self/cls argument 

703 

704 Example:: 

705 

706 >>> format_argspec_plus(lambda self, a, b, c=3, **d: 123) 

707 {'grouped_args': '(self, a, b, c=3, **d)', 

708 'self_arg': 'self', 

709 'apply_kw': '(self, a, b, c=c, **d)', 

710 'apply_pos': '(self, a, b, c, **d)'} 

711 

712 """ 

713 if callable(fn): 

714 spec = compat.inspect_getfullargspec(fn) 

715 else: 

716 spec = fn 

717 

718 args = compat.inspect_formatargspec(*spec) 

719 

720 apply_pos = compat.inspect_formatargspec( 

721 spec[0], spec[1], spec[2], None, spec[4] 

722 ) 

723 

724 if spec[0]: 

725 self_arg = spec[0][0] 

726 

727 apply_pos_proxied = compat.inspect_formatargspec( 

728 spec[0][1:], spec[1], spec[2], None, spec[4] 

729 ) 

730 

731 elif spec[1]: 

732 # I'm not sure what this is 

733 self_arg = "%s[0]" % spec[1] 

734 

735 apply_pos_proxied = apply_pos 

736 else: 

737 self_arg = None 

738 apply_pos_proxied = apply_pos 

739 

740 num_defaults = 0 

741 if spec[3]: 

742 num_defaults += len(cast(Tuple[Any], spec[3])) 

743 if spec[4]: 

744 num_defaults += len(spec[4]) 

745 

746 name_args = spec[0] + spec[4] 

747 

748 defaulted_vals: Union[List[str], Tuple[()]] 

749 

750 if num_defaults: 

751 defaulted_vals = name_args[0 - num_defaults :] 

752 else: 

753 defaulted_vals = () 

754 

755 apply_kw = compat.inspect_formatargspec( 

756 name_args, 

757 spec[1], 

758 spec[2], 

759 defaulted_vals, 

760 formatvalue=lambda x: "=" + str(x), 

761 ) 

762 

763 if spec[0]: 

764 apply_kw_proxied = compat.inspect_formatargspec( 

765 name_args[1:], 

766 spec[1], 

767 spec[2], 

768 defaulted_vals, 

769 formatvalue=lambda x: "=" + str(x), 

770 ) 

771 else: 

772 apply_kw_proxied = apply_kw 

773 

774 if grouped: 

775 return dict( 

776 grouped_args=args, 

777 self_arg=self_arg, 

778 apply_pos=apply_pos, 

779 apply_kw=apply_kw, 

780 apply_pos_proxied=apply_pos_proxied, 

781 apply_kw_proxied=apply_kw_proxied, 

782 ) 

783 else: 

784 return dict( 

785 grouped_args=args, 

786 self_arg=self_arg, 

787 apply_pos=apply_pos[1:-1], 

788 apply_kw=apply_kw[1:-1], 

789 apply_pos_proxied=apply_pos_proxied[1:-1], 

790 apply_kw_proxied=apply_kw_proxied[1:-1], 

791 ) 

792 

793 

794def format_argspec_init(method, grouped=True): 

795 """format_argspec_plus with considerations for typical __init__ methods 

796 

797 Wraps format_argspec_plus with error handling strategies for typical 

798 __init__ cases: 

799 

800 .. sourcecode:: text 

801 

802 object.__init__ -> (self) 

803 other unreflectable (usually C) -> (self, *args, **kwargs) 

804 

805 """ 

806 if method is object.__init__: 

807 grouped_args = "(self)" 

808 args = "(self)" if grouped else "self" 

809 proxied = "()" if grouped else "" 

810 else: 

811 try: 

812 return format_argspec_plus(method, grouped=grouped) 

813 except TypeError: 

814 grouped_args = "(self, *args, **kwargs)" 

815 args = grouped_args if grouped else "self, *args, **kwargs" 

816 proxied = "(*args, **kwargs)" if grouped else "*args, **kwargs" 

817 return dict( 

818 self_arg="self", 

819 grouped_args=grouped_args, 

820 apply_pos=args, 

821 apply_kw=args, 

822 apply_pos_proxied=proxied, 

823 apply_kw_proxied=proxied, 

824 ) 

825 

826 

827def create_proxy_methods( 

828 target_cls: Type[Any], 

829 target_cls_sphinx_name: str, 

830 proxy_cls_sphinx_name: str, 

831 classmethods: Sequence[str] = (), 

832 methods: Sequence[str] = (), 

833 attributes: Sequence[str] = (), 

834 use_intermediate_variable: Sequence[str] = (), 

835) -> Callable[[_T], _T]: 

836 """A class decorator indicating attributes should refer to a proxy 

837 class. 

838 

839 This decorator is now a "marker" that does nothing at runtime. Instead, 

840 it is consumed by the tools/generate_proxy_methods.py script to 

841 statically generate proxy methods and attributes that are fully 

842 recognized by typing tools such as mypy. 

843 

844 """ 

845 

846 def decorate(cls): 

847 return cls 

848 

849 return decorate 

850 

851 

852def getargspec_init(method): 

853 """inspect.getargspec with considerations for typical __init__ methods 

854 

855 Wraps inspect.getargspec with error handling for typical __init__ cases: 

856 

857 .. sourcecode:: text 

858 

859 object.__init__ -> (self) 

860 other unreflectable (usually C) -> (self, *args, **kwargs) 

861 

862 """ 

863 try: 

864 return compat.inspect_getfullargspec(method) 

865 except TypeError: 

866 if method is object.__init__: 

867 return (["self"], None, None, None) 

868 else: 

869 return (["self"], "args", "kwargs", None) 

870 

871 

872def unbound_method_to_callable(func_or_cls): 

873 """Adjust the incoming callable such that a 'self' argument is not 

874 required. 

875 

876 """ 

877 

878 if isinstance(func_or_cls, types.MethodType) and not func_or_cls.__self__: 

879 return func_or_cls.__func__ 

880 else: 

881 return func_or_cls 

882 

883 

884class GenericRepr: 

885 """Encapsulates the logic for creating a generic __repr__() string. 

886 

887 This class allows for the repr structure to be created, then modified 

888 (e.g., changing the class name), before being rendered as a string. 

889 

890 .. versionadded:: 2.1 

891 """ 

892 

893 __slots__ = ( 

894 "_obj", 

895 "_additional_kw", 

896 "_to_inspect", 

897 "_omit_kwarg", 

898 "_class_name", 

899 ) 

900 

901 _obj: Any 

902 _additional_kw: Sequence[Tuple[str, Any]] 

903 _to_inspect: List[object] 

904 _omit_kwarg: Sequence[str] 

905 _class_name: Optional[str] 

906 

907 def __init__( 

908 self, 

909 obj: Any, 

910 additional_kw: Sequence[Tuple[str, Any]] = (), 

911 to_inspect: Optional[Union[object, List[object]]] = None, 

912 omit_kwarg: Sequence[str] = (), 

913 ): 

914 """Create a GenericRepr object. 

915 

916 :param obj: The object being repr'd 

917 :param additional_kw: Additional keyword arguments to check for in 

918 the repr, as a sequence of 2-tuples of (name, default_value) 

919 :param to_inspect: One or more objects whose __init__ signature 

920 should be inspected. If not provided, defaults to [obj]. 

921 :param omit_kwarg: Sequence of keyword argument names to omit from 

922 the repr output 

923 """ 

924 self._obj = obj 

925 self._additional_kw = additional_kw 

926 self._to_inspect = ( 

927 [obj] if to_inspect is None else _collections.to_list(to_inspect) 

928 ) 

929 self._omit_kwarg = omit_kwarg 

930 self._class_name = None 

931 

932 def set_class_name(self, class_name: str) -> GenericRepr: 

933 """Set the class name to be used in the repr. 

934 

935 By default, the class name is taken from obj.__class__.__name__. 

936 This method allows it to be overridden. 

937 

938 :param class_name: The class name to use 

939 :return: self, for method chaining 

940 """ 

941 self._class_name = class_name 

942 return self 

943 

944 def __str__(self) -> str: 

945 """Produce the __repr__() string based on the configured parameters.""" 

946 obj = self._obj 

947 to_inspect = self._to_inspect 

948 additional_kw = self._additional_kw 

949 omit_kwarg = self._omit_kwarg 

950 

951 missing = object() 

952 

953 pos_args = [] 

954 kw_args: _collections.OrderedDict[str, Any] = ( 

955 _collections.OrderedDict() 

956 ) 

957 vargs = None 

958 for i, insp in enumerate(to_inspect): 

959 try: 

960 spec = compat.inspect_getfullargspec(insp.__init__) # type: ignore[misc] # noqa: E501 

961 except TypeError: 

962 continue 

963 else: 

964 default_len = len(spec.defaults) if spec.defaults else 0 

965 if i == 0: 

966 if spec.varargs: 

967 vargs = spec.varargs 

968 if default_len: 

969 pos_args.extend(spec.args[1:-default_len]) 

970 else: 

971 pos_args.extend(spec.args[1:]) 

972 else: 

973 kw_args.update( 

974 [(arg, missing) for arg in spec.args[1:-default_len]] 

975 ) 

976 

977 if default_len: 

978 assert spec.defaults 

979 kw_args.update( 

980 [ 

981 (arg, default) 

982 for arg, default in zip( 

983 spec.args[-default_len:], spec.defaults 

984 ) 

985 ] 

986 ) 

987 output: List[str] = [] 

988 

989 output.extend(repr(getattr(obj, arg, None)) for arg in pos_args) 

990 

991 if vargs is not None and hasattr(obj, vargs): 

992 output.extend([repr(val) for val in getattr(obj, vargs)]) 

993 

994 for arg, defval in kw_args.items(): 

995 if arg in omit_kwarg: 

996 continue 

997 try: 

998 val = getattr(obj, arg, missing) 

999 if val is not missing and val != defval: 

1000 output.append("%s=%r" % (arg, val)) 

1001 except Exception: 

1002 pass 

1003 

1004 if additional_kw: 

1005 for arg, defval in additional_kw: 

1006 try: 

1007 val = getattr(obj, arg, missing) 

1008 if val is not missing and val != defval: 

1009 output.append("%s=%r" % (arg, val)) 

1010 except Exception: 

1011 pass 

1012 

1013 class_name = ( 

1014 self._class_name 

1015 if self._class_name is not None 

1016 else obj.__class__.__name__ 

1017 ) 

1018 return "%s(%s)" % (class_name, ", ".join(output)) 

1019 

1020 

1021def generic_repr( 

1022 obj: Any, 

1023 additional_kw: Sequence[Tuple[str, Any]] = (), 

1024 to_inspect: Optional[Union[object, List[object]]] = None, 

1025 omit_kwarg: Sequence[str] = (), 

1026) -> str: 

1027 """Produce a __repr__() based on direct association of the __init__() 

1028 specification vs. same-named attributes present. 

1029 

1030 """ 

1031 return str( 

1032 GenericRepr( 

1033 obj, 

1034 additional_kw=additional_kw, 

1035 to_inspect=to_inspect, 

1036 omit_kwarg=omit_kwarg, 

1037 ) 

1038 ) 

1039 

1040 

1041def class_hierarchy(cls): 

1042 """Return an unordered sequence of all classes related to cls. 

1043 

1044 Traverses diamond hierarchies. 

1045 

1046 Fibs slightly: subclasses of builtin types are not returned. Thus 

1047 class_hierarchy(class A(object)) returns (A, object), not A plus every 

1048 class systemwide that derives from object. 

1049 

1050 """ 

1051 

1052 hier = {cls} 

1053 process = list(cls.__mro__) 

1054 while process: 

1055 c = process.pop() 

1056 bases = (_ for _ in c.__bases__ if _ not in hier) 

1057 

1058 for b in bases: 

1059 process.append(b) 

1060 hier.add(b) 

1061 

1062 if c.__module__ == "builtins" or not hasattr(c, "__subclasses__"): 

1063 continue 

1064 

1065 for s in [ 

1066 _ 

1067 for _ in ( 

1068 c.__subclasses__() 

1069 if not issubclass(c, type) 

1070 else c.__subclasses__(c) 

1071 ) 

1072 if _ not in hier 

1073 ]: 

1074 process.append(s) 

1075 hier.add(s) 

1076 return list(hier) 

1077 

1078 

1079def iterate_attributes(cls): 

1080 """iterate all the keys and attributes associated 

1081 with a class, without using getattr(). 

1082 

1083 Does not use getattr() so that class-sensitive 

1084 descriptors (i.e. property.__get__()) are not called. 

1085 

1086 """ 

1087 keys = dir(cls) 

1088 for key in keys: 

1089 for c in cls.__mro__: 

1090 if key in c.__dict__: 

1091 yield (key, c.__dict__[key]) 

1092 break 

1093 

1094 

1095def monkeypatch_proxied_specials( 

1096 into_cls, 

1097 from_cls, 

1098 skip=None, 

1099 only=None, 

1100 name="self.proxy", 

1101 from_instance=None, 

1102): 

1103 """Automates delegation of __specials__ for a proxying type.""" 

1104 

1105 if only: 

1106 dunders = only 

1107 else: 

1108 if skip is None: 

1109 skip = ( 

1110 "__slots__", 

1111 "__del__", 

1112 "__getattribute__", 

1113 "__metaclass__", 

1114 "__getstate__", 

1115 "__setstate__", 

1116 ) 

1117 dunders = [ 

1118 m 

1119 for m in dir(from_cls) 

1120 if ( 

1121 m.startswith("__") 

1122 and m.endswith("__") 

1123 and not hasattr(into_cls, m) 

1124 and m not in skip 

1125 ) 

1126 ] 

1127 

1128 for method in dunders: 

1129 try: 

1130 maybe_fn = getattr(from_cls, method) 

1131 if not hasattr(maybe_fn, "__call__"): 

1132 continue 

1133 maybe_fn = getattr(maybe_fn, "__func__", maybe_fn) 

1134 fn = cast(types.FunctionType, maybe_fn) 

1135 

1136 except AttributeError: 

1137 continue 

1138 try: 

1139 spec = compat.inspect_getfullargspec(fn) 

1140 fn_args = compat.inspect_formatargspec(spec[0]) 

1141 d_args = compat.inspect_formatargspec(spec[0][1:]) 

1142 except TypeError: 

1143 fn_args = "(self, *args, **kw)" 

1144 d_args = "(*args, **kw)" 

1145 

1146 py = ( 

1147 "def %(method)s%(fn_args)s: " 

1148 "return %(name)s.%(method)s%(d_args)s" % locals() 

1149 ) 

1150 

1151 env: Dict[str, types.FunctionType] = ( 

1152 from_instance is not None and {name: from_instance} or {} 

1153 ) 

1154 # the generated source is derived entirely from from_cls and the method 

1155 # name, and into_cls is a throwaway per-descriptor class whose 

1156 # __qualname__ is the same for every one of them, so this is keyed on 

1157 # from_cls rather than into_cls 

1158 proxied = exec_code_in_env( 

1159 py, 

1160 env, 

1161 method, 

1162 f"{method} proxying to " 

1163 f"{from_cls.__module__}.{from_cls.__qualname__}", 

1164 ) 

1165 try: 

1166 proxied.__defaults__ = fn.__defaults__ 

1167 except AttributeError: 

1168 pass 

1169 setattr(into_cls, method, proxied) 

1170 

1171 

1172def methods_equivalent(meth1, meth2): 

1173 """Return True if the two methods are the same implementation.""" 

1174 

1175 return getattr(meth1, "__func__", meth1) is getattr( 

1176 meth2, "__func__", meth2 

1177 ) 

1178 

1179 

1180def as_interface(obj, cls=None, methods=None, required=None): 

1181 """Ensure basic interface compliance for an instance or dict of callables. 

1182 

1183 Checks that ``obj`` implements public methods of ``cls`` or has members 

1184 listed in ``methods``. If ``required`` is not supplied, implementing at 

1185 least one interface method is sufficient. Methods present on ``obj`` that 

1186 are not in the interface are ignored. 

1187 

1188 If ``obj`` is a dict and ``dict`` does not meet the interface 

1189 requirements, the keys of the dictionary are inspected. Keys present in 

1190 ``obj`` that are not in the interface will raise TypeErrors. 

1191 

1192 Raises TypeError if ``obj`` does not meet the interface criteria. 

1193 

1194 In all passing cases, an object with callable members is returned. In the 

1195 simple case, ``obj`` is returned as-is; if dict processing kicks in then 

1196 an anonymous class is returned. 

1197 

1198 obj 

1199 A type, instance, or dictionary of callables. 

1200 cls 

1201 Optional, a type. All public methods of cls are considered the 

1202 interface. An ``obj`` instance of cls will always pass, ignoring 

1203 ``required``.. 

1204 methods 

1205 Optional, a sequence of method names to consider as the interface. 

1206 required 

1207 Optional, a sequence of mandatory implementations. If omitted, an 

1208 ``obj`` that provides at least one interface method is considered 

1209 sufficient. As a convenience, required may be a type, in which case 

1210 all public methods of the type are required. 

1211 

1212 """ 

1213 if not cls and not methods: 

1214 raise TypeError("a class or collection of method names are required") 

1215 

1216 if isinstance(cls, type) and isinstance(obj, cls): 

1217 return obj 

1218 

1219 interface = set(methods or [m for m in dir(cls) if not m.startswith("_")]) 

1220 implemented = set(dir(obj)) 

1221 

1222 complies = operator.ge 

1223 if isinstance(required, type): 

1224 required = interface 

1225 elif not required: 

1226 required = set() 

1227 complies = operator.gt 

1228 else: 

1229 required = set(required) 

1230 

1231 if complies(implemented.intersection(interface), required): 

1232 return obj 

1233 

1234 # No dict duck typing here. 

1235 if not isinstance(obj, dict): 

1236 qualifier = complies is operator.gt and "any of" or "all of" 

1237 raise TypeError( 

1238 "%r does not implement %s: %s" 

1239 % (obj, qualifier, ", ".join(interface)) 

1240 ) 

1241 

1242 class AnonymousInterface: 

1243 """A callable-holding shell.""" 

1244 

1245 if cls: 

1246 AnonymousInterface.__name__ = "Anonymous" + cls.__name__ 

1247 found = set() 

1248 

1249 for method, impl in dictlike_iteritems(obj): 

1250 if method not in interface: 

1251 raise TypeError("%r: unknown in this interface" % method) 

1252 if not callable(impl): 

1253 raise TypeError("%r=%r is not callable" % (method, impl)) 

1254 setattr(AnonymousInterface, method, staticmethod(impl)) 

1255 found.add(method) 

1256 

1257 if complies(found, required): 

1258 return AnonymousInterface 

1259 

1260 raise TypeError( 

1261 "dictionary does not contain required keys %s" 

1262 % ", ".join(required - found) 

1263 ) 

1264 

1265 

1266_GFD = TypeVar("_GFD", bound="generic_fn_descriptor[Any]") 

1267 

1268 

1269class generic_fn_descriptor(Generic[_T_co]): 

1270 """Descriptor which proxies a function when the attribute is not 

1271 present in dict 

1272 

1273 This superclass is organized in a particular way with "memoized" and 

1274 "non-memoized" implementation classes that are hidden from type checkers, 

1275 as Mypy seems to not be able to handle seeing multiple kinds of descriptor 

1276 classes used for the same attribute. 

1277 

1278 """ 

1279 

1280 fget: Callable[..., _T_co] 

1281 __doc__: Optional[str] 

1282 __name__: str 

1283 

1284 def __init__(self, fget: Callable[..., _T_co], doc: Optional[str] = None): 

1285 self.fget = fget 

1286 self.__doc__ = doc or fget.__doc__ 

1287 self.__name__ = fget.__name__ 

1288 

1289 @overload 

1290 def __get__(self: _GFD, obj: None, cls: Any) -> _GFD: ... 

1291 

1292 @overload 

1293 def __get__(self, obj: object, cls: Any) -> _T_co: ... 

1294 

1295 def __get__(self: _GFD, obj: Any, cls: Any) -> Union[_GFD, _T_co]: 

1296 raise NotImplementedError() 

1297 

1298 if TYPE_CHECKING: 

1299 

1300 def __set__(self, instance: Any, value: Any) -> None: ... 

1301 

1302 def __delete__(self, instance: Any) -> None: ... 

1303 

1304 def _reset(self, obj: Any) -> None: 

1305 raise NotImplementedError() 

1306 

1307 @classmethod 

1308 def reset(cls, obj: Any, name: str) -> None: 

1309 raise NotImplementedError() 

1310 

1311 

1312class _non_memoized_property(generic_fn_descriptor[_T_co]): 

1313 """a plain descriptor that proxies a function. 

1314 

1315 primary rationale is to provide a plain attribute that's 

1316 compatible with memoized_property which is also recognized as equivalent 

1317 by mypy. 

1318 

1319 """ 

1320 

1321 if not TYPE_CHECKING: 

1322 

1323 def __get__(self, obj, cls): 

1324 if obj is None: 

1325 return self 

1326 return self.fget(obj) 

1327 

1328 

1329class _memoized_property(generic_fn_descriptor[_T_co]): 

1330 """A read-only @property that is only evaluated once.""" 

1331 

1332 if not TYPE_CHECKING: 

1333 

1334 def __get__(self, obj, cls): 

1335 if obj is None: 

1336 return self 

1337 obj.__dict__[self.__name__] = result = self.fget(obj) 

1338 return result 

1339 

1340 def _reset(self, obj): 

1341 _memoized_property.reset(obj, self.__name__) 

1342 

1343 @classmethod 

1344 def reset(cls, obj, name): 

1345 obj.__dict__.pop(name, None) 

1346 

1347 

1348# despite many attempts to get Mypy to recognize an overridden descriptor 

1349# where one is memoized and the other isn't, there seems to be no reliable 

1350# way other than completely deceiving the type checker into thinking there 

1351# is just one single descriptor type everywhere. Otherwise, if a superclass 

1352# has non-memoized and subclass has memoized, that requires 

1353# "class memoized(non_memoized)". but then if a superclass has memoized and 

1354# superclass has non-memoized, the class hierarchy of the descriptors 

1355# would need to be reversed; "class non_memoized(memoized)". so there's no 

1356# way to achieve this. 

1357# additional issues, RO properties: 

1358# https://github.com/python/mypy/issues/12440 

1359if TYPE_CHECKING: 

1360 # allow memoized and non-memoized to be freely mixed by having them 

1361 # be the same class 

1362 memoized_property = generic_fn_descriptor 

1363 non_memoized_property = generic_fn_descriptor 

1364 

1365 # for read only situations, mypy only sees @property as read only. 

1366 # read only is needed when a subtype specializes the return type 

1367 # of a property, meaning assignment needs to be disallowed 

1368 ro_memoized_property = property 

1369 ro_non_memoized_property = property 

1370 

1371else: 

1372 memoized_property = ro_memoized_property = _memoized_property 

1373 non_memoized_property = ro_non_memoized_property = _non_memoized_property 

1374 

1375 

1376def memoized_instancemethod(fn: _F) -> _F: 

1377 """Decorate a method memoize its return value. 

1378 

1379 Best applied to no-arg methods: memoization is not sensitive to 

1380 argument values, and will always return the same value even when 

1381 called with different arguments. 

1382 

1383 """ 

1384 

1385 def oneshot(self, *args, **kw): 

1386 result = fn(self, *args, **kw) 

1387 

1388 def memo(*a, **kw): 

1389 return result 

1390 

1391 memo.__name__ = fn.__name__ 

1392 memo.__doc__ = fn.__doc__ 

1393 self.__dict__[fn.__name__] = memo 

1394 return result 

1395 

1396 return update_wrapper(oneshot, fn) # type: ignore[return-value] 

1397 

1398 

1399class HasMemoized: 

1400 """A mixin class that maintains the names of memoized elements in a 

1401 collection for easy cache clearing, generative, etc. 

1402 

1403 """ 

1404 

1405 if not TYPE_CHECKING: 

1406 # support classes that want to have __slots__ with an explicit 

1407 # slot for __dict__. not sure if that requires base __slots__ here. 

1408 __slots__ = () 

1409 

1410 _memoized_keys: FrozenSet[str] = frozenset() 

1411 

1412 def _reset_memoizations(self) -> None: 

1413 for elem in self._memoized_keys: 

1414 self.__dict__.pop(elem, None) 

1415 

1416 def _assert_no_memoizations(self) -> None: 

1417 for elem in self._memoized_keys: 

1418 assert elem not in self.__dict__ 

1419 

1420 def _set_memoized_attribute(self, key: str, value: Any) -> None: 

1421 self.__dict__[key] = value 

1422 self._memoized_keys |= {key} 

1423 

1424 class memoized_attribute(memoized_property[_T]): 

1425 """A read-only @property that is only evaluated once. 

1426 

1427 :meta private: 

1428 

1429 """ 

1430 

1431 fget: Callable[..., _T] 

1432 __doc__: Optional[str] 

1433 __name__: str 

1434 

1435 def __init__(self, fget: Callable[..., _T], doc: Optional[str] = None): 

1436 self.fget = fget 

1437 self.__doc__ = doc or fget.__doc__ 

1438 self.__name__ = fget.__name__ 

1439 

1440 @overload 

1441 def __get__(self: _MA, obj: None, cls: Any) -> _MA: ... 

1442 

1443 @overload 

1444 def __get__(self, obj: Any, cls: Any) -> _T: ... 

1445 

1446 def __get__(self, obj, cls): 

1447 if obj is None: 

1448 return self 

1449 obj.__dict__[self.__name__] = result = self.fget(obj) 

1450 obj._memoized_keys |= {self.__name__} 

1451 return result 

1452 

1453 @classmethod 

1454 def memoized_instancemethod(cls, fn: _F) -> _F: 

1455 """Decorate a method memoize its return value. 

1456 

1457 :meta private: 

1458 

1459 """ 

1460 

1461 def oneshot(self: Any, *args: Any, **kw: Any) -> Any: 

1462 result = fn(self, *args, **kw) 

1463 

1464 def memo(*a, **kw): 

1465 return result 

1466 

1467 memo.__name__ = fn.__name__ 

1468 memo.__doc__ = fn.__doc__ 

1469 self.__dict__[fn.__name__] = memo 

1470 self._memoized_keys |= {fn.__name__} 

1471 return result 

1472 

1473 return update_wrapper(oneshot, fn) # type: ignore[return-value] 

1474 

1475 

1476if TYPE_CHECKING: 

1477 HasMemoized_ro_memoized_attribute = property 

1478else: 

1479 HasMemoized_ro_memoized_attribute = HasMemoized.memoized_attribute 

1480 

1481 

1482class MemoizedSlots: 

1483 """Apply memoized items to an object using a __getattr__ scheme. 

1484 

1485 This allows the functionality of memoized_property and 

1486 memoized_instancemethod to be available to a class using __slots__. 

1487 

1488 The memoized get is not threadsafe under freethreading and the 

1489 creator method may in extremely rare cases be called more than once. 

1490 

1491 """ 

1492 

1493 __slots__ = () 

1494 

1495 def _fallback_getattr(self, key): 

1496 raise AttributeError(key) 

1497 

1498 def __getattr__(self, key: str) -> Any: 

1499 if key.startswith("_memoized_attr_") or key.startswith( 

1500 "_memoized_method_" 

1501 ): 

1502 raise AttributeError(key) 

1503 # to avoid recursion errors when interacting with other __getattr__ 

1504 # schemes that refer to this one, when testing for memoized method 

1505 # look at __class__ only rather than going into __getattr__ again. 

1506 elif hasattr(self.__class__, f"_memoized_attr_{key}"): 

1507 value = getattr(self, f"_memoized_attr_{key}")() 

1508 setattr(self, key, value) 

1509 return value 

1510 elif hasattr(self.__class__, f"_memoized_method_{key}"): 

1511 meth = getattr(self, f"_memoized_method_{key}") 

1512 

1513 def oneshot(*args, **kw): 

1514 result = meth(*args, **kw) 

1515 

1516 def memo(*a, **kw): 

1517 return result 

1518 

1519 memo.__name__ = meth.__name__ 

1520 memo.__doc__ = meth.__doc__ 

1521 setattr(self, key, memo) 

1522 return result 

1523 

1524 oneshot.__doc__ = meth.__doc__ 

1525 return oneshot 

1526 else: 

1527 return self._fallback_getattr(key) 

1528 

1529 

1530# from paste.deploy.converters 

1531def asbool(obj: Any) -> bool: 

1532 if isinstance(obj, str): 

1533 obj = obj.strip().lower() 

1534 if obj in ["true", "yes", "on", "y", "t", "1"]: 

1535 return True 

1536 elif obj in ["false", "no", "off", "n", "f", "0"]: 

1537 return False 

1538 else: 

1539 raise ValueError("String is not true/false: %r" % obj) 

1540 return bool(obj) 

1541 

1542 

1543def bool_or_str(*text: str) -> Callable[[str], Union[str, bool]]: 

1544 """Return a callable that will evaluate a string as 

1545 boolean, or one of a set of "alternate" string values. 

1546 

1547 """ 

1548 

1549 def bool_or_value(obj: str) -> Union[str, bool]: 

1550 if obj in text: 

1551 return obj 

1552 else: 

1553 return asbool(obj) 

1554 

1555 return bool_or_value 

1556 

1557 

1558def asint(value: Any) -> Optional[int]: 

1559 """Coerce to integer.""" 

1560 

1561 if value is None: 

1562 return value 

1563 return int(value) 

1564 

1565 

1566def coerce_kw_type( 

1567 kw: Dict[str, Any], 

1568 key: str, 

1569 type_: Type[Any], 

1570 flexi_bool: bool = True, 

1571 dest: Optional[Dict[str, Any]] = None, 

1572) -> None: 

1573 r"""If 'key' is present in dict 'kw', coerce its value to type 'type\_' if 

1574 necessary. If 'flexi_bool' is True, the string '0' is considered false 

1575 when coercing to boolean. 

1576 """ 

1577 

1578 if dest is None: 

1579 dest = kw 

1580 

1581 if ( 

1582 key in kw 

1583 and (not isinstance(type_, type) or not isinstance(kw[key], type_)) 

1584 and kw[key] is not None 

1585 ): 

1586 if type_ is bool and flexi_bool: 

1587 dest[key] = asbool(kw[key]) 

1588 else: 

1589 dest[key] = type_(kw[key]) 

1590 

1591 

1592def constructor_key(obj: Any, cls: Type[Any]) -> Tuple[Any, ...]: 

1593 """Produce a tuple structure that is cacheable using the __dict__ of 

1594 obj to retrieve values 

1595 

1596 """ 

1597 names = get_cls_kwargs(cls) 

1598 return (cls,) + tuple( 

1599 (k, obj.__dict__[k]) for k in names if k in obj.__dict__ 

1600 ) 

1601 

1602 

1603def constructor_copy(obj: _T, cls: Type[_T], *args: Any, **kw: Any) -> _T: 

1604 """Instantiate cls using the __dict__ of obj as constructor arguments. 

1605 

1606 Uses inspect to match the named arguments of ``cls``. 

1607 

1608 """ 

1609 

1610 names = get_cls_kwargs(cls) 

1611 kw.update( 

1612 (k, obj.__dict__[k]) for k in names.difference(kw) if k in obj.__dict__ 

1613 ) 

1614 return cls(*args, **kw) 

1615 

1616 

1617def counter() -> Callable[[], int]: 

1618 """Return a threadsafe counter function.""" 

1619 

1620 lock = threading.Lock() 

1621 counter = itertools.count(1) 

1622 

1623 # avoid the 2to3 "next" transformation... 

1624 def _next(): 

1625 with lock: 

1626 return next(counter) 

1627 

1628 return _next 

1629 

1630 

1631def duck_type_collection( 

1632 specimen: Any, default: Optional[Type[Any]] = None 

1633) -> Optional[Type[Any]]: 

1634 """Given an instance or class, guess if it is or is acting as one of 

1635 the basic collection types: list, set and dict. If the __emulates__ 

1636 property is present, return that preferentially. 

1637 """ 

1638 

1639 if hasattr(specimen, "__emulates__"): 

1640 # canonicalize set vs sets.Set to a standard: the builtin set 

1641 if specimen.__emulates__ is not None and issubclass( 

1642 specimen.__emulates__, set 

1643 ): 

1644 return set 

1645 else: 

1646 return specimen.__emulates__ # type: ignore[no-any-return] 

1647 

1648 isa = issubclass if isinstance(specimen, type) else isinstance 

1649 if isa(specimen, list): 

1650 return list 

1651 elif isa(specimen, set): 

1652 return set 

1653 elif isa(specimen, dict): 

1654 return dict 

1655 

1656 if hasattr(specimen, "append"): 

1657 return list 

1658 elif hasattr(specimen, "add"): 

1659 return set 

1660 elif hasattr(specimen, "set"): 

1661 return dict 

1662 else: 

1663 return default 

1664 

1665 

1666def assert_arg_type( 

1667 arg: Any, argtype: Union[Tuple[Type[Any], ...], Type[Any]], name: str 

1668) -> Any: 

1669 if isinstance(arg, argtype): 

1670 return arg 

1671 else: 

1672 if isinstance(argtype, tuple): 

1673 raise exc.ArgumentError( 

1674 "Argument '%s' is expected to be one of type %s, got '%s'" 

1675 % (name, " or ".join("'%s'" % a for a in argtype), type(arg)) 

1676 ) 

1677 else: 

1678 raise exc.ArgumentError( 

1679 "Argument '%s' is expected to be of type '%s', got '%s'" 

1680 % (name, argtype, type(arg)) 

1681 ) 

1682 

1683 

1684def dictlike_iteritems(dictlike): 

1685 """Return a (key, value) iterator for almost any dict-like object.""" 

1686 

1687 if hasattr(dictlike, "items"): 

1688 return list(dictlike.items()) 

1689 

1690 getter = getattr(dictlike, "__getitem__", getattr(dictlike, "get", None)) 

1691 if getter is None: 

1692 raise TypeError("Object '%r' is not dict-like" % dictlike) 

1693 

1694 if hasattr(dictlike, "iterkeys"): 

1695 

1696 def iterator(): 

1697 for key in dictlike.iterkeys(): 

1698 assert getter is not None 

1699 yield key, getter(key) 

1700 

1701 return iterator() 

1702 elif hasattr(dictlike, "keys"): 

1703 return iter((key, getter(key)) for key in dictlike.keys()) 

1704 else: 

1705 raise TypeError("Object '%r' is not dict-like" % dictlike) 

1706 

1707 

1708class classproperty(property): 

1709 """A decorator that behaves like @property except that operates 

1710 on classes rather than instances. 

1711 

1712 The decorator is currently special when using the declarative 

1713 module, but note that the 

1714 :class:`~.sqlalchemy.ext.declarative.declared_attr` 

1715 decorator should be used for this purpose with declarative. 

1716 

1717 """ 

1718 

1719 fget: Callable[[Any], Any] 

1720 

1721 def __init__(self, fget: Callable[[Any], Any], *arg: Any, **kw: Any): 

1722 super().__init__(fget, *arg, **kw) 

1723 self.__doc__ = fget.__doc__ 

1724 

1725 def __get__(self, obj: Any, cls: Optional[type] = None) -> Any: 

1726 return self.fget(cls) 

1727 

1728 

1729class hybridproperty(Generic[_T]): 

1730 def __init__(self, func: Callable[..., _T]): 

1731 self.func = func 

1732 self.clslevel = func 

1733 

1734 def __get__(self, instance: Any, owner: Any) -> _T: 

1735 if instance is None: 

1736 clsval = self.clslevel(owner) 

1737 return clsval 

1738 else: 

1739 return self.func(instance) 

1740 

1741 def classlevel(self, func: Callable[..., Any]) -> hybridproperty[_T]: 

1742 self.clslevel = func 

1743 return self 

1744 

1745 

1746class rw_hybridproperty(Generic[_T]): 

1747 def __init__(self, func: Callable[..., _T]): 

1748 self.func = func 

1749 self.clslevel = func 

1750 self.setfn: Optional[Callable[..., Any]] = None 

1751 

1752 def __get__(self, instance: Any, owner: Any) -> _T: 

1753 if instance is None: 

1754 clsval = self.clslevel(owner) 

1755 return clsval 

1756 else: 

1757 return self.func(instance) 

1758 

1759 def __set__(self, instance: Any, value: Any) -> None: 

1760 assert self.setfn is not None 

1761 self.setfn(instance, value) 

1762 

1763 def setter(self, func: Callable[..., Any]) -> rw_hybridproperty[_T]: 

1764 self.setfn = func 

1765 return self 

1766 

1767 def classlevel(self, func: Callable[..., Any]) -> rw_hybridproperty[_T]: 

1768 self.clslevel = func 

1769 return self 

1770 

1771 

1772class hybridmethod(Generic[_T]): 

1773 """Decorate a function as cls- or instance- level.""" 

1774 

1775 def __init__(self, func: Callable[..., _T]): 

1776 self.func = self.__func__ = func 

1777 self.clslevel = func 

1778 

1779 def __get__(self, instance: Any, owner: Any) -> Callable[..., _T]: 

1780 if instance is None: 

1781 return self.clslevel.__get__( # type: ignore[no-any-return] 

1782 owner, owner.__class__ 

1783 ) 

1784 else: 

1785 return self.func.__get__( # type: ignore[no-any-return] 

1786 instance, owner 

1787 ) 

1788 

1789 def classlevel(self, func: Callable[..., Any]) -> hybridmethod[_T]: 

1790 self.clslevel = func 

1791 return self 

1792 

1793 

1794class symbol(int): 

1795 """A constant symbol. 

1796 

1797 >>> symbol("foo") is symbol("foo") 

1798 True 

1799 >>> symbol("foo") 

1800 <symbol 'foo> 

1801 

1802 A slight refinement of the MAGICCOOKIE=object() pattern. The primary 

1803 advantage of symbol() is its repr(). They are also singletons. 

1804 

1805 Repeated calls of symbol('name') will all return the same instance. 

1806 

1807 """ 

1808 

1809 name: str 

1810 

1811 symbols: Dict[str, symbol] = {} 

1812 _lock = threading.Lock() 

1813 

1814 def __new__( 

1815 cls, 

1816 name: str, 

1817 doc: Optional[str] = None, 

1818 canonical: Optional[int] = None, 

1819 ) -> symbol: 

1820 with cls._lock: 

1821 sym = cls.symbols.get(name) 

1822 if sym is None: 

1823 assert isinstance(name, str) 

1824 if canonical is None: 

1825 canonical = hash(name) 

1826 sym = int.__new__(symbol, canonical) 

1827 sym.name = name 

1828 if doc: 

1829 sym.__doc__ = doc 

1830 

1831 # NOTE: we should ultimately get rid of this global thing, 

1832 # however, currently it is to support pickling. The best 

1833 # change would be when we are on py3.11 at a minimum, we 

1834 # switch to stdlib enum.IntFlag. 

1835 cls.symbols[name] = sym 

1836 else: 

1837 if canonical and canonical != sym: 

1838 raise TypeError( 

1839 f"Can't replace canonical symbol for {name!r} " 

1840 f"with new int value {canonical}" 

1841 ) 

1842 return sym 

1843 

1844 def __reduce__(self): 

1845 return symbol, (self.name, "x", int(self)) 

1846 

1847 def __str__(self): 

1848 return repr(self) 

1849 

1850 def __repr__(self): 

1851 return f"symbol({self.name!r})" 

1852 

1853 

1854class _IntFlagMeta(type): 

1855 def __init__( 

1856 cls, 

1857 classname: str, 

1858 bases: Tuple[Type[Any], ...], 

1859 dict_: Dict[str, Any], 

1860 **kw: Any, 

1861 ) -> None: 

1862 items: List[symbol] 

1863 cls._items = items = [] 

1864 for k, v in dict_.items(): 

1865 if re.match(r"^__.*__$", k): 

1866 continue 

1867 if isinstance(v, int): 

1868 sym = symbol(k, canonical=v) 

1869 elif not k.startswith("_"): 

1870 raise TypeError("Expected integer values for IntFlag") 

1871 else: 

1872 continue 

1873 setattr(cls, k, sym) 

1874 items.append(sym) 

1875 

1876 cls.__members__ = _collections.immutabledict( 

1877 {sym.name: sym for sym in items} 

1878 ) 

1879 

1880 def __iter__(self) -> Iterator[symbol]: 

1881 raise NotImplementedError( 

1882 "iter not implemented to ensure compatibility with " 

1883 "Python 3.11 IntFlag. Please use __members__. See " 

1884 "https://github.com/python/cpython/issues/99304" 

1885 ) 

1886 

1887 

1888class _FastIntFlag(metaclass=_IntFlagMeta): 

1889 """An 'IntFlag' copycat that isn't slow when performing bitwise 

1890 operations. 

1891 

1892 the ``FastIntFlag`` class will return ``enum.IntFlag`` under TYPE_CHECKING 

1893 and ``_FastIntFlag`` otherwise. 

1894 

1895 """ 

1896 

1897 

1898if TYPE_CHECKING: 

1899 from enum import IntFlag 

1900 

1901 FastIntFlag = IntFlag 

1902else: 

1903 FastIntFlag = _FastIntFlag 

1904 

1905 

1906_E = TypeVar("_E", bound=enum.Enum) 

1907 

1908 

1909def parse_user_argument_for_enum( 

1910 arg: Any, 

1911 choices: Dict[_E, List[Any]], 

1912 name: str, 

1913 resolve_symbol_names: bool = False, 

1914) -> Optional[_E]: 

1915 """Given a user parameter, parse the parameter into a chosen value 

1916 from a list of choice objects, typically Enum values. 

1917 

1918 The user argument can be a string name that matches the name of a 

1919 symbol, or the symbol object itself, or any number of alternate choices 

1920 such as True/False/ None etc. 

1921 

1922 :param arg: the user argument. 

1923 :param choices: dictionary of enum values to lists of possible 

1924 entries for each. 

1925 :param name: name of the argument. Used in an :class:`.ArgumentError` 

1926 that is raised if the parameter doesn't match any available argument. 

1927 

1928 """ 

1929 for enum_value, choice in choices.items(): 

1930 if arg is enum_value: 

1931 return enum_value 

1932 elif resolve_symbol_names and arg == enum_value.name: 

1933 return enum_value 

1934 elif arg in choice: 

1935 return enum_value 

1936 

1937 if arg is None: 

1938 return None 

1939 

1940 raise exc.ArgumentError(f"Invalid value for '{name}': {arg!r}") 

1941 

1942 

1943_creation_order = 1 

1944 

1945 

1946def set_creation_order(instance: Any) -> None: 

1947 """Assign a '_creation_order' sequence to the given instance. 

1948 

1949 This allows multiple instances to be sorted in order of creation 

1950 (typically within a single thread; the counter is not particularly 

1951 threadsafe). 

1952 

1953 """ 

1954 global _creation_order 

1955 instance._creation_order = _creation_order 

1956 _creation_order += 1 

1957 

1958 

1959def warn_exception(func: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: 

1960 """executes the given function, catches all exceptions and converts to 

1961 a warning. 

1962 

1963 """ 

1964 try: 

1965 return func(*args, **kwargs) 

1966 except Exception: 

1967 warn("%s('%s') ignored" % sys.exc_info()[0:2]) 

1968 

1969 

1970def ellipses_string(value, len_=25): 

1971 try: 

1972 if len(value) > len_: 

1973 return "%s..." % value[0:len_] 

1974 else: 

1975 return value 

1976 except TypeError: 

1977 return value 

1978 

1979 

1980class _hash_limit_string(str): 

1981 """A string subclass that can only be hashed on a maximum amount 

1982 of unique values. 

1983 

1984 This is used for warnings so that we can send out parameterized warnings 

1985 without the __warningregistry__ of the module, or the non-overridable 

1986 "once" registry within warnings.py, overloading memory, 

1987 

1988 

1989 """ 

1990 

1991 _hash: int 

1992 

1993 def __new__( 

1994 cls, value: str, num: int, args: Sequence[Any] 

1995 ) -> _hash_limit_string: 

1996 interpolated = (value % args) + ( 

1997 " (this warning may be suppressed after %d occurrences)" % num 

1998 ) 

1999 self = super().__new__(cls, interpolated) 

2000 self._hash = hash("%s_%d" % (value, hash(interpolated) % num)) 

2001 return self 

2002 

2003 def __hash__(self) -> int: 

2004 return self._hash 

2005 

2006 def __eq__(self, other: Any) -> bool: 

2007 return hash(self) == hash(other) 

2008 

2009 

2010def warn(msg: str, code: Optional[str] = None) -> None: 

2011 """Issue a warning. 

2012 

2013 If msg is a string, :class:`.exc.SAWarning` is used as 

2014 the category. 

2015 

2016 """ 

2017 if code: 

2018 _warnings_warn(exc.SAWarning(msg, code=code)) 

2019 else: 

2020 _warnings_warn(msg, exc.SAWarning) 

2021 

2022 

2023def warn_limited(msg: str, args: Sequence[Any]) -> None: 

2024 """Issue a warning with a parameterized string, limiting the number 

2025 of registrations. 

2026 

2027 """ 

2028 if args: 

2029 msg = _hash_limit_string(msg, 10, args) 

2030 _warnings_warn(msg, exc.SAWarning) 

2031 

2032 

2033_warning_tags: Dict[CodeType, Tuple[str, Type[Warning]]] = {} 

2034 

2035 

2036def tag_method_for_warnings( 

2037 message: str, category: Type[Warning] 

2038) -> Callable[[_F], _F]: 

2039 def go(fn): 

2040 _warning_tags[fn.__code__] = (message, category) 

2041 return fn 

2042 

2043 return go 

2044 

2045 

2046_not_sa_pattern = re.compile(r"^(?:sqlalchemy\.(?!testing)|alembic\.)") 

2047 

2048 

2049def _warnings_warn( 

2050 message: Union[str, Warning], 

2051 category: Optional[Type[Warning]] = None, 

2052 stacklevel: int = 2, 

2053) -> None: 

2054 

2055 if category is None and isinstance(message, Warning): 

2056 category = type(message) 

2057 

2058 # adjust the given stacklevel to be outside of SQLAlchemy 

2059 try: 

2060 frame = sys._getframe(stacklevel) 

2061 except ValueError: 

2062 # being called from less than 3 (or given) stacklevels, weird, 

2063 # but don't crash 

2064 stacklevel = 0 

2065 except: 

2066 # _getframe() doesn't work, weird interpreter issue, weird, 

2067 # ok, but don't crash 

2068 stacklevel = 0 

2069 else: 

2070 stacklevel_found = warning_tag_found = False 

2071 while frame is not None: 

2072 # using __name__ here requires that we have __name__ in the 

2073 # __globals__ of the decorated string functions we make also. 

2074 # we generate this using {"__name__": fn.__module__} 

2075 if not stacklevel_found and not re.match( 

2076 _not_sa_pattern, frame.f_globals.get("__name__", "") 

2077 ): 

2078 # stop incrementing stack level if an out-of-SQLA line 

2079 # were found. 

2080 stacklevel_found = True 

2081 

2082 # however, for the warning tag thing, we have to keep 

2083 # scanning up the whole traceback 

2084 

2085 if frame.f_code in _warning_tags: 

2086 warning_tag_found = True 

2087 _suffix, _category = _warning_tags[frame.f_code] 

2088 category = category or _category 

2089 message = f"{message} ({_suffix})" 

2090 

2091 frame = frame.f_back # type: ignore[assignment] 

2092 

2093 if not stacklevel_found: 

2094 stacklevel += 1 

2095 elif stacklevel_found and warning_tag_found: 

2096 break 

2097 

2098 if category is not None: 

2099 warnings.warn(message, category, stacklevel=stacklevel + 1) 

2100 else: 

2101 warnings.warn(message, stacklevel=stacklevel + 1) 

2102 

2103 

2104def only_once( 

2105 fn: Callable[..., _T], retry_on_exception: bool 

2106) -> Callable[..., Optional[_T]]: 

2107 """Decorate the given function to be a no-op after it is called exactly 

2108 once.""" 

2109 

2110 once = [fn] 

2111 

2112 def go(*arg: Any, **kw: Any) -> Optional[_T]: 

2113 # strong reference fn so that it isn't garbage collected, 

2114 # which interferes with the event system's expectations 

2115 strong_fn = fn # noqa 

2116 if once: 

2117 once_fn = once.pop() 

2118 try: 

2119 return once_fn(*arg, **kw) 

2120 except: 

2121 if retry_on_exception: 

2122 once.insert(0, once_fn) 

2123 raise 

2124 

2125 return None 

2126 

2127 return go 

2128 

2129 

2130_SQLA_RE = re.compile(r"sqlalchemy/([a-z_]+/){0,2}[a-z_]+\.py") 

2131_UNITTEST_RE = re.compile(r"unit(?:2|test2?/)") 

2132 

2133 

2134def chop_traceback( 

2135 tb: List[str], 

2136 exclude_prefix: re.Pattern[str] = _UNITTEST_RE, 

2137 exclude_suffix: re.Pattern[str] = _SQLA_RE, 

2138) -> List[str]: 

2139 """Chop extraneous lines off beginning and end of a traceback. 

2140 

2141 :param tb: 

2142 a list of traceback lines as returned by ``traceback.format_stack()`` 

2143 

2144 :param exclude_prefix: 

2145 a regular expression object matching lines to skip at beginning of 

2146 ``tb`` 

2147 

2148 :param exclude_suffix: 

2149 a regular expression object matching lines to skip at end of ``tb`` 

2150 """ 

2151 start = 0 

2152 end = len(tb) - 1 

2153 while start <= end and exclude_prefix.search(tb[start]): 

2154 start += 1 

2155 while start <= end and exclude_suffix.search(tb[end]): 

2156 end -= 1 

2157 return tb[start : end + 1] 

2158 

2159 

2160def attrsetter(attrname): 

2161 code = "def set(obj, value): obj.%s = value" % attrname 

2162 env = locals().copy() 

2163 exec(code, env) 

2164 return env["set"] 

2165 

2166 

2167dunders_re = re.compile("^__.+__$") 

2168 

2169 

2170class TypingOnly: 

2171 """A mixin class that marks a class as 'typing only', meaning it has 

2172 absolutely no methods, attributes, or runtime functionality whatsoever. 

2173 

2174 """ 

2175 

2176 __slots__ = () 

2177 

2178 def __init_subclass__(cls, **kw: Any) -> None: 

2179 if TypingOnly in cls.__bases__: 

2180 remaining = { 

2181 name for name in cls.__dict__ if not dunders_re.match(name) 

2182 } 

2183 if remaining: 

2184 raise AssertionError( 

2185 f"Class {cls} directly inherits TypingOnly but has " 

2186 f"additional attributes {remaining}." 

2187 ) 

2188 super().__init_subclass__(**kw) 

2189 

2190 

2191class EnsureKWArg: 

2192 r"""Apply translation of functions to accept \**kw arguments if they 

2193 don't already. 

2194 

2195 Used to ensure cross-compatibility with third party legacy code, for things 

2196 like compiler visit methods that need to accept ``**kw`` arguments, 

2197 but may have been copied from old code that didn't accept them. 

2198 

2199 """ 

2200 

2201 ensure_kwarg: str 

2202 """a regular expression that indicates method names for which the method 

2203 should accept ``**kw`` arguments. 

2204 

2205 The class will scan for methods matching the name template and decorate 

2206 them if necessary to ensure ``**kw`` parameters are accepted. 

2207 

2208 """ 

2209 

2210 def __init_subclass__(cls) -> None: 

2211 fn_reg = cls.ensure_kwarg 

2212 clsdict = cls.__dict__ 

2213 if fn_reg: 

2214 for key in clsdict: 

2215 m = re.match(fn_reg, key) 

2216 if m: 

2217 fn = clsdict[key] 

2218 spec = compat.inspect_getfullargspec(fn) 

2219 if not spec.varkw: 

2220 wrapped = cls._wrap_w_kw(fn) 

2221 setattr(cls, key, wrapped) 

2222 super().__init_subclass__() 

2223 

2224 @classmethod 

2225 def _wrap_w_kw(cls, fn: Callable[..., Any]) -> Callable[..., Any]: 

2226 def wrap(*arg: Any, **kw: Any) -> Any: 

2227 return fn(*arg) 

2228 

2229 return update_wrapper(wrap, fn) 

2230 

2231 

2232def wrap_callable(wrapper, fn): 

2233 """Augment functools.update_wrapper() to work with objects with 

2234 a ``__call__()`` method. 

2235 

2236 :param fn: 

2237 object with __call__ method 

2238 

2239 """ 

2240 if hasattr(fn, "__name__"): 

2241 return update_wrapper(wrapper, fn) 

2242 else: 

2243 _f = wrapper 

2244 _f.__name__ = fn.__class__.__name__ 

2245 if hasattr(fn, "__module__"): 

2246 _f.__module__ = fn.__module__ 

2247 

2248 if hasattr(fn.__call__, "__doc__") and fn.__call__.__doc__: 

2249 _f.__doc__ = fn.__call__.__doc__ 

2250 elif fn.__doc__: 

2251 _f.__doc__ = fn.__doc__ 

2252 

2253 return _f 

2254 

2255 

2256def find_matching_paren(text: str, start: int = 0) -> Optional[int]: 

2257 """Return the index of the ``)`` that matches the ``(`` at ``start``. 

2258 

2259 The walk skips single-quoted (``'...'``) and double-quoted (``"..."``) 

2260 string literals, so parentheses inside string literals do not affect 

2261 the depth counter. ``''`` and ``""`` are treated as escaped quotes 

2262 inside their respective contexts, matching PostgreSQL/SQLite literal 

2263 conventions. 

2264 

2265 Returns ``None`` if the opening parenthesis is never closed 

2266 (unbalanced). The character at ``text[start]`` must be ``(``. 

2267 

2268 Note for SQLite use, SQLite also supports MySQL backtick-style quotes as 

2269 well as SQL Server bracket style quotes; the latter has different escaping 

2270 behaviors. A follow-up patch could add support for these two additional 

2271 styles (consider using an enum like QuotingStyle.DOUBLE | 

2272 QuotingStyle.BRACKET, etc.) 

2273 

2274 E.g.:: 

2275 

2276 >>> find_matching_paren("(a + b)") 

2277 6 

2278 >>> find_matching_paren("((a)(b))") 

2279 7 

2280 >>> find_matching_paren("(a = '(' AND b = ')')") 

2281 20 

2282 

2283 """ 

2284 assert text[start] == "(", "start index must point at an open paren" 

2285 

2286 depth = 0 

2287 in_single = False 

2288 in_double = False 

2289 n = len(text) 

2290 i = start 

2291 while i < n: 

2292 ch = text[i] 

2293 if in_single: 

2294 if ch == "'": 

2295 if i + 1 < n and text[i + 1] == "'": 

2296 i += 2 

2297 continue 

2298 in_single = False 

2299 elif in_double: 

2300 if ch == '"': 

2301 if i + 1 < n and text[i + 1] == '"': 

2302 i += 2 

2303 continue 

2304 in_double = False 

2305 elif ch == "'": 

2306 in_single = True 

2307 elif ch == '"': 

2308 in_double = True 

2309 elif ch == "(": 

2310 depth += 1 

2311 elif ch == ")": 

2312 depth -= 1 

2313 if depth == 0: 

2314 return i 

2315 i += 1 

2316 return None 

2317 

2318 

2319def strip_outer_parens(text: str) -> str: 

2320 """Remove one layer of outer parentheses from ``text`` if they wrap the 

2321 entire (stripped) string. 

2322 

2323 Whitespace is preserved if the parentheses do not wrap the whole 

2324 expression. String literals are honored via :func:`find_matching_paren`, 

2325 so ``"(a = '(' AND b = ')')"`` correctly strips to 

2326 ``"a = '(' AND b = ')'"`` rather than being interpreted as two separate 

2327 paren groups. 

2328 

2329 E.g.:: 

2330 

2331 >>> strip_outer_parens("(a IS NOT NULL)") 

2332 'a IS NOT NULL' 

2333 >>> strip_outer_parens("(a) AND (b)") 

2334 '(a) AND (b)' 

2335 >>> strip_outer_parens("a NOT NULL") 

2336 'a NOT NULL' 

2337 

2338 """ 

2339 stripped = text.strip() 

2340 lstripped = len(stripped) 

2341 if lstripped < 2 or stripped[0] != "(" or stripped[-1] != ")": 

2342 return text 

2343 close = find_matching_paren(stripped, 0) 

2344 if close is not None and close == lstripped - 1: 

2345 return stripped[1:-1] 

2346 return text 

2347 

2348 

2349def quoted_token_parser(value): 

2350 """Parse a dotted identifier with accommodation for quoted names. 

2351 

2352 Includes support for SQL-style double quotes as a literal character. 

2353 

2354 E.g.:: 

2355 

2356 >>> quoted_token_parser("name") 

2357 ["name"] 

2358 >>> quoted_token_parser("schema.name") 

2359 ["schema", "name"] 

2360 >>> quoted_token_parser('"Schema"."Name"') 

2361 ['Schema', 'Name'] 

2362 >>> quoted_token_parser('"Schema"."Name""Foo"') 

2363 ['Schema', 'Name""Foo'] 

2364 

2365 """ 

2366 

2367 if '"' not in value: 

2368 return value.split(".") 

2369 

2370 # 0 = outside of quotes 

2371 # 1 = inside of quotes 

2372 state = 0 

2373 result: List[List[str]] = [[]] 

2374 idx = 0 

2375 lv = len(value) 

2376 while idx < lv: 

2377 char = value[idx] 

2378 if char == '"': 

2379 if state == 1 and idx < lv - 1 and value[idx + 1] == '"': 

2380 result[-1].append('"') 

2381 idx += 1 

2382 else: 

2383 state ^= 1 

2384 elif char == "." and state == 0: 

2385 result.append([]) 

2386 else: 

2387 result[-1].append(char) 

2388 idx += 1 

2389 

2390 return ["".join(token) for token in result] 

2391 

2392 

2393def add_parameter_text(params: Any, text: str) -> Callable[[_F], _F]: 

2394 params = _collections.to_list(params) 

2395 

2396 def decorate(fn): 

2397 doc = fn.__doc__ is not None and fn.__doc__ or "" 

2398 if doc: 

2399 doc = inject_param_text(doc, {param: text for param in params}) 

2400 fn.__doc__ = doc 

2401 return fn 

2402 

2403 return decorate 

2404 

2405 

2406def _dedent_docstring(text: str) -> str: 

2407 split_text = text.split("\n", 1) 

2408 if len(split_text) == 1: 

2409 return text 

2410 else: 

2411 firstline, remaining = split_text 

2412 if not firstline.startswith(" "): 

2413 return firstline + "\n" + textwrap.dedent(remaining) 

2414 else: 

2415 return textwrap.dedent(text) 

2416 

2417 

2418def inject_docstring_text( 

2419 given_doctext: Optional[str], injecttext: str, pos: int 

2420) -> str: 

2421 doctext: str = _dedent_docstring(given_doctext or "") 

2422 lines = doctext.split("\n") 

2423 if len(lines) == 1: 

2424 lines.append("") 

2425 injectlines = textwrap.dedent(injecttext).split("\n") 

2426 if injectlines[0]: 

2427 injectlines.insert(0, "") 

2428 

2429 blanks = [num for num, line in enumerate(lines) if not line.strip()] 

2430 blanks.insert(0, 0) 

2431 

2432 inject_pos = blanks[min(pos, len(blanks) - 1)] 

2433 

2434 lines = lines[0:inject_pos] + injectlines + lines[inject_pos:] 

2435 return "\n".join(lines) 

2436 

2437 

2438_param_reg = re.compile(r"(\s+):param (.+?):") 

2439 

2440 

2441def inject_param_text(doctext: str, inject_params: Dict[str, str]) -> str: 

2442 doclines = collections.deque(doctext.splitlines()) 

2443 lines = [] 

2444 

2445 # TODO: this is not working for params like ":param case_sensitive=True:" 

2446 

2447 to_inject = None 

2448 while doclines: 

2449 line = doclines.popleft() 

2450 

2451 m = _param_reg.match(line) 

2452 

2453 if to_inject is None: 

2454 if m: 

2455 param = m.group(2).lstrip("*") 

2456 if param in inject_params: 

2457 # default indent to that of :param: plus one 

2458 indent = " " * len(m.group(1)) + " " 

2459 

2460 # but if the next line has text, use that line's 

2461 # indentation 

2462 if doclines: 

2463 m2 = re.match(r"(\s+)\S", doclines[0]) 

2464 if m2: 

2465 indent = " " * len(m2.group(1)) 

2466 

2467 to_inject = indent + inject_params[param] 

2468 elif m: 

2469 lines.extend(["\n", to_inject, "\n"]) 

2470 to_inject = None 

2471 elif not line.rstrip(): 

2472 lines.extend([line, to_inject, "\n"]) 

2473 to_inject = None 

2474 elif line.endswith("::"): 

2475 # TODO: this still won't cover if the code example itself has 

2476 # blank lines in it, need to detect those via indentation. 

2477 lines.extend([line, doclines.popleft()]) 

2478 continue 

2479 lines.append(line) 

2480 

2481 return "\n".join(lines) 

2482 

2483 

2484def repr_tuple_names(names: List[str]) -> Optional[str]: 

2485 """Trims a list of strings from the middle and return a string of up to 

2486 four elements. Strings greater than 11 characters will be truncated""" 

2487 if len(names) == 0: 

2488 return None 

2489 flag = len(names) <= 4 

2490 names = names[0:4] if flag else names[0:3] + names[-1:] 

2491 res = ["%s.." % name[:11] if len(name) > 11 else name for name in names] 

2492 if flag: 

2493 return ", ".join(res) 

2494 else: 

2495 return "%s, ..., %s" % (", ".join(res[0:3]), res[-1]) 

2496 

2497 

2498def has_compiled_ext(raise_=False): 

2499 from ._has_cython import HAS_CYEXTENSION 

2500 

2501 if HAS_CYEXTENSION: 

2502 return True 

2503 elif raise_: 

2504 raise ImportError( 

2505 "cython extensions were expected to be installed, " 

2506 "but are not present" 

2507 ) 

2508 else: 

2509 return False 

2510 

2511 

2512def load_uncompiled_module(module: _M) -> _M: 

2513 """Load the non-compied version of a module that is also 

2514 compiled with cython. 

2515 """ 

2516 full_name = module.__name__ 

2517 assert module.__spec__ 

2518 parent_name = module.__spec__.parent 

2519 assert parent_name 

2520 parent_module = sys.modules[parent_name] 

2521 assert parent_module.__spec__ 

2522 package_path = parent_module.__spec__.origin 

2523 assert package_path and package_path.endswith("__init__.py") 

2524 

2525 name = full_name.split(".")[-1] 

2526 module_path = package_path.replace("__init__.py", f"{name}.py") 

2527 

2528 py_spec = importlib.util.spec_from_file_location(full_name, module_path) 

2529 assert py_spec 

2530 py_module = importlib.util.module_from_spec(py_spec) 

2531 assert py_spec.loader 

2532 py_spec.loader.exec_module(py_module) 

2533 return cast(_M, py_module) 

2534 

2535 

2536_pre_release_normalize = { 

2537 "a": "a", 

2538 "alpha": "a", 

2539 "b": "b", 

2540 "beta": "b", 

2541 "c": "rc", 

2542 "pre": "rc", 

2543 "preview": "rc", 

2544 "rc": "rc", 

2545} 

2546 

2547_version_string_re = re.compile( 

2548 r""" 

2549 \s* 

2550 (?:[a-z][a-z0-9]*[-_])? # ignored prefix, "py3-" 

2551 v? 

2552 (?P<release>\d+(?:\.\d+)*) 

2553 (?: # pre-release 

2554 [-_.]? 

2555 (?P<pre_l>alpha|beta|preview|pre|rc|a|b|c) 

2556 [-_.]? 

2557 (?P<pre_n>\d+)? 

2558 )? 

2559 (?: # post-release 

2560 [-_.]? 

2561 (?P<post_l>post|rev|r) 

2562 [-_.]? 

2563 (?P<post_n>\d+)? 

2564 )? 

2565 (?: # developmental release 

2566 [-_.]? 

2567 (?P<dev_l>dev) 

2568 [-_.]? 

2569 (?P<dev_n>\d+)? 

2570 )? 

2571 """, 

2572 re.X | re.I, 

2573) 

2574 

2575_VersionSortKey = Tuple[ 

2576 Tuple[int, ...], 

2577 Tuple[int, str, int], 

2578 Tuple[int, int], 

2579 Tuple[int, int], 

2580] 

2581 

2582 

2583def _version_sort_key( 

2584 release: Tuple[int, ...], 

2585 pre: Optional[Tuple[str, int]], 

2586 post: Optional[int], 

2587 dev: Optional[int], 

2588) -> _VersionSortKey: 

2589 if pre is None and post is None and dev is not None: 

2590 # a dev release with no other qualifiers precedes every 

2591 # pre-release of the same release number 

2592 pre_key = (-1, "", 0) 

2593 elif pre is None: 

2594 pre_key = (1, "", 0) 

2595 else: 

2596 pre_key = (0, pre[0], pre[1]) 

2597 

2598 return ( 

2599 release, 

2600 pre_key, 

2601 (0, 0) if post is None else (1, post), 

2602 (1, 0) if dev is None else (0, dev), 

2603 ) 

2604 

2605 

2606def _version_comparison( 

2607 op: Callable[[Any, Any], bool], 

2608) -> Callable[[VersionInfo, Any], Any]: 

2609 """Build one of :class:`.VersionInfo`'s comparison methods. 

2610 

2611 Comparison takes place against the sort key rather than the tuple 

2612 itself, so that pre-release and similar qualifiers are taken into 

2613 account. A plain tuple is interpreted as the release segment of a 

2614 final release; anything else is not comparable. 

2615 

2616 """ 

2617 

2618 def compare(self: VersionInfo, other: Any) -> Any: 

2619 if isinstance(other, VersionInfo): 

2620 other_key = other._sort_key 

2621 elif isinstance(other, tuple): 

2622 other_key = _version_sort_key(other, None, None, None) 

2623 else: 

2624 return NotImplemented 

2625 return op(self._sort_key, other_key) 

2626 

2627 return compare 

2628 

2629 

2630class VersionInfo(Tuple[int, ...]): 

2631 """A version number, as a tuple of integers. 

2632 

2633 :class:`.VersionInfo` is a ``tuple`` subclass consisting of the 

2634 numeric "release" segment of a version only, e.g. ``2.0.0rc1`` 

2635 is the tuple ``(2, 0, 0)``. Ordering however takes any 

2636 pre-release, post-release and developmental qualifiers into account 

2637 as described by :pep:`440`, so that ``2.0.0rc1`` compares as less than 

2638 ``2.0.0``, including when compared against a plain tuple such as 

2639 ``(2, 0, 0)``. 

2640 

2641 Plain tuples are interpreted as final releases when compared against 

2642 a :class:`.VersionInfo`. 

2643 

2644 .. versionadded:: 2.1 

2645 

2646 """ 

2647 

2648 string: Optional[str] 

2649 """the string from which this version was parsed, if any.""" 

2650 

2651 pre: Optional[Tuple[str, int]] 

2652 """normalized pre-release qualifier, e.g. ``("rc", 1)``.""" 

2653 

2654 post: Optional[int] 

2655 """post-release number, if any.""" 

2656 

2657 dev: Optional[int] 

2658 """developmental release number, if any.""" 

2659 

2660 _sort_key: _VersionSortKey 

2661 

2662 def __new__( 

2663 cls, 

2664 release: Sequence[int] = (), 

2665 *, 

2666 string: Optional[str] = None, 

2667 pre: Optional[Tuple[str, int]] = None, 

2668 post: Optional[int] = None, 

2669 dev: Optional[int] = None, 

2670 ) -> VersionInfo: 

2671 # __new__ is needed as the release segment has to be passed to 

2672 # tuple.__new__(); the remaining state is set up in __init__ 

2673 return tuple.__new__(cls, release) 

2674 

2675 def __init__( 

2676 self, 

2677 release: Sequence[int] = (), 

2678 *, 

2679 string: Optional[str] = None, 

2680 pre: Optional[Tuple[str, int]] = None, 

2681 post: Optional[int] = None, 

2682 dev: Optional[int] = None, 

2683 ): 

2684 self.string = string 

2685 self.pre = pre 

2686 self.post = post 

2687 self.dev = dev 

2688 self._sort_key = _version_sort_key(tuple(self), pre, post, dev) 

2689 

2690 def __repr__(self) -> str: 

2691 if self.string is not None: 

2692 return f"VersionInfo({tuple(self)!r}, string={self.string!r})" 

2693 else: 

2694 return f"VersionInfo({tuple(self)!r})" 

2695 

2696 def __str__(self) -> str: 

2697 if self.string is not None: 

2698 return self.string 

2699 else: 

2700 return ".".join(str(num) for num in self) 

2701 

2702 # every comparison has to be stated explicitly; ``tuple`` implements 

2703 # all six of them, so ``functools.total_ordering`` fills in nothing 

2704 # here and the ones left out would silently compare as plain tuples 

2705 __eq__ = _version_comparison(operator.eq) 

2706 __ne__ = _version_comparison(operator.ne) 

2707 __lt__ = _version_comparison(operator.lt) 

2708 __le__ = _version_comparison(operator.le) 

2709 __gt__ = _version_comparison(operator.gt) 

2710 __ge__ = _version_comparison(operator.ge) 

2711 

2712 def __hash__(self) -> int: 

2713 return hash(self._sort_key) 

2714 

2715 

2716def parse_version_string(version: Optional[str]) -> VersionInfo: 

2717 """Parse a DBAPI version string into a :class:`.VersionInfo`. 

2718 

2719 Leading characters that are not part of the version itself are 

2720 ignored, as are trailing characters following the version, so that 

2721 strings such as ``"py3-4.0.19-beta4"`` and 

2722 ``"2.9.10 (dt dec pq3 ext lo64)"`` parse correctly. 

2723 

2724 An empty :class:`.VersionInfo` is returned if no version number can be 

2725 located at all. 

2726 

2727 Parsing is deliberately more tolerant than that of :pep:`440`, which 

2728 the version strings published by DBAPIs frequently do not conform to; 

2729 a strict implementation such as that of the ``packaging`` library 

2730 rejects each of the above outright. 

2731 

2732 .. versionadded:: 2.1 

2733 

2734 """ 

2735 

2736 if not version: 

2737 return VersionInfo((), string=version) 

2738 

2739 m = _version_string_re.match(version) 

2740 if m is None: 

2741 return VersionInfo((), string=version) 

2742 

2743 release = tuple(int(x) for x in m.group("release").split(".")) 

2744 

2745 pre_l = m.group("pre_l") 

2746 pre: Optional[Tuple[str, int]] 

2747 if pre_l is not None: 

2748 pre = ( 

2749 _pre_release_normalize[pre_l.lower()], 

2750 int(m.group("pre_n") or 0), 

2751 ) 

2752 else: 

2753 pre = None 

2754 

2755 return VersionInfo( 

2756 release, 

2757 string=version, 

2758 pre=pre, 

2759 post=( 

2760 int(m.group("post_n") or 0) 

2761 if m.group("post_l") is not None 

2762 else None 

2763 ), 

2764 dev=( 

2765 int(m.group("dev_n") or 0) 

2766 if m.group("dev_l") is not None 

2767 else None 

2768 ), 

2769 ) 

2770 

2771 

2772def parse_version_from_metadata(distribution: str) -> VersionInfo: 

2773 """Return the version of an installed distribution as a 

2774 :class:`.VersionInfo`. 

2775 

2776 This is intended for use by dialects whose DBAPI module does not 

2777 itself publish a version number, such as ``asyncmy``. As the 

2778 distribution name is not necessarily the same as the module name, and 

2779 the installed distribution is not necessarily the module that was 

2780 imported, this should not be used when the DBAPI module provides a 

2781 version of its own. 

2782 

2783 An empty :class:`.VersionInfo` is returned if the distribution is not 

2784 installed. 

2785 

2786 .. versionadded:: 2.1 

2787 

2788 """ 

2789 

2790 try: 

2791 version = importlib.metadata.version(distribution) 

2792 except importlib.metadata.PackageNotFoundError: 

2793 return VersionInfo() 

2794 else: 

2795 return parse_version_string(version) 

2796 

2797 

2798class _Missing(enum.Enum): 

2799 Missing = enum.auto() 

2800 

2801 

2802Missing = _Missing.Missing 

2803MissingOr = Union[_T, Literal[_Missing.Missing]]