codekingpro/portable-devtools
114k
1# util/_collections.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"""Collection classes and helpers."""
10from __future__ import annotations
11
12import operator
13import threading
14import types
15import typing
16from typing import Any
17from typing import Callable
18from typing import cast
19from typing import Dict
20from typing import FrozenSet
21from typing import Generic
22from typing import Iterable
23from typing import Iterator
24from typing import List
25from typing import Mapping
26from typing import NoReturn
27from typing import Optional
28from typing import overload
29from typing import Sequence
30from typing import Set
31from typing import Tuple
32from typing import TypeVar
33from typing import Union
34from typing import ValuesView
35import weakref
36
37from ._has_cy import HAS_CYEXTENSION
38from .typing import is_non_string_iterable
39from .typing import Literal
40from .typing import Protocol
41
42if typing.TYPE_CHECKING or not HAS_CYEXTENSION:
43 from ._py_collections import immutabledict as immutabledict
44 from ._py_collections import IdentitySet as IdentitySet
45 from ._py_collections import ReadOnlyContainer as ReadOnlyContainer
46 from ._py_collections import ImmutableDictBase as ImmutableDictBase
47 from ._py_collections import OrderedSet as OrderedSet
48 from ._py_collections import unique_list as unique_list
49else:
50 from sqlalchemy.cyextension.immutabledict import (
51 ReadOnlyContainer as ReadOnlyContainer,
52 )
53 from sqlalchemy.cyextension.immutabledict import (
54 ImmutableDictBase as ImmutableDictBase,
55 )
56 from sqlalchemy.cyextension.immutabledict import (
57 immutabledict as immutabledict,
58 )
59 from sqlalchemy.cyextension.collections import IdentitySet as IdentitySet
60 from sqlalchemy.cyextension.collections import OrderedSet as OrderedSet
61 from sqlalchemy.cyextension.collections import ( # noqa
62 unique_list as unique_list,
63 )
64
65
66_T = TypeVar("_T", bound=Any)
67_KT = TypeVar("_KT", bound=Any)
68_VT = TypeVar("_VT", bound=Any)
69_T_co = TypeVar("_T_co", covariant=True)
70
71EMPTY_SET: FrozenSet[Any] = frozenset()
72NONE_SET: FrozenSet[Any] = frozenset([None])
73
74
75def merge_lists_w_ordering(a: List[Any], b: List[Any]) -> List[Any]:
76 """merge two lists, maintaining ordering as much as possible.
77
78 this is to reconcile vars(cls) with cls.__annotations__.
79
80 Example::
81
82 >>> a = ['__tablename__', 'id', 'x', 'created_at']
83 >>> b = ['id', 'name', 'data', 'y', 'created_at']
84 >>> merge_lists_w_ordering(a, b)
85 ['__tablename__', 'id', 'name', 'data', 'y', 'x', 'created_at']
86
87 This is not necessarily the ordering that things had on the class,
88 in this case the class is::
89
90 class User(Base):
91 __tablename__ = "users"
92
93 id: Mapped[int] = mapped_column(primary_key=True)
94 name: Mapped[str]
95 data: Mapped[Optional[str]]
96 x = Column(Integer)
97 y: Mapped[int]
98 created_at: Mapped[datetime.datetime] = mapped_column()
99
100 But things are *mostly* ordered.
101
102 The algorithm could also be done by creating a partial ordering for
103 all items in both lists and then using topological_sort(), but that
104 is too much overhead.
105
106 Background on how I came up with this is at:
107 https://gist.github.com/zzzeek/89de958cf0803d148e74861bd682ebae
108
109 """
110 overlap = set(a).intersection(b)
111
112 result = []
113
114 current, other = iter(a), iter(b)
115
116 while True:
117 for element in current:
118 if element in overlap:
119 overlap.discard(element)
120 other, current = current, other
121 break
122
123 result.append(element)
124 else:
125 result.extend(other)
126 break
127
128 return result
129
130
131def coerce_to_immutabledict(d: Mapping[_KT, _VT]) -> immutabledict[_KT, _VT]:
132 if not d:
133 return EMPTY_DICT
134 elif isinstance(d, immutabledict):
135 return d
136 else:
137 return immutabledict(d)
138
139
140EMPTY_DICT: immutabledict[Any, Any] = immutabledict()
141
142
143class FacadeDict(ImmutableDictBase[_KT, _VT]):
144 """A dictionary that is not publicly mutable."""
145
146 def __new__(cls, *args: Any) -> FacadeDict[Any, Any]:
147 new = ImmutableDictBase.__new__(cls)
148 return new
149
150 def copy(self) -> NoReturn:
151 raise NotImplementedError(
152 "an immutabledict shouldn't need to be copied. use dict(d) "
153 "if you need a mutable dictionary."
154 )
155
156 def __reduce__(self) -> Any:
157 return FacadeDict, (dict(self),)
158
159 def _insert_item(self, key: _KT, value: _VT) -> None:
160 """insert an item into the dictionary directly."""
161 dict.__setitem__(self, key, value)
162
163 def __repr__(self) -> str:
164 return "FacadeDict(%s)" % dict.__repr__(self)
165
166
167_DT = TypeVar("_DT", bound=Any)
168
169_F = TypeVar("_F", bound=Any)
170
171
172class Properties(Generic[_T]):
173 """Provide a __getattr__/__setattr__ interface over a dict."""
174
175 __slots__ = ("_data",)
176
177 _data: Dict[str, _T]
178
179 def __init__(self, data: Dict[str, _T]):
180 object.__setattr__(self, "_data", data)
181
182 def __len__(self) -> int:
183 return len(self._data)
184
185 def __iter__(self) -> Iterator[_T]:
186 return iter(list(self._data.values()))
187
188 def __dir__(self) -> List[str]:
189 return dir(super()) + [str(k) for k in self._data.keys()]
190
191 def __add__(self, other: Properties[_F]) -> List[Union[_T, _F]]:
192 return list(self) + list(other)
193
194 def __setitem__(self, key: str, obj: _T) -> None:
195 self._data[key] = obj
196
197 def __getitem__(self, key: str) -> _T:
198 return self._data[key]
199
200 def __delitem__(self, key: str) -> None:
201 del self._data[key]
202
203 def __setattr__(self, key: str, obj: _T) -> None:
204 self._data[key] = obj
205
206 def __getstate__(self) -> Dict[str, Any]:
207 return {"_data": self._data}
208
209 def __setstate__(self, state: Dict[str, Any]) -> None:
210 object.__setattr__(self, "_data", state["_data"])
211
212 def __getattr__(self, key: str) -> _T:
213 try:
214 return self._data[key]
215 except KeyError:
216 raise AttributeError(key)
217
218 def __contains__(self, key: str) -> bool:
219 return key in self._data
220
221 def as_readonly(self) -> ReadOnlyProperties[_T]:
222 """Return an immutable proxy for this :class:`.Properties`."""
223
224 return ReadOnlyProperties(self._data)
225
226 def update(self, value: Dict[str, _T]) -> None:
227 self._data.update(value)
228
229 @overload
230 def get(self, key: str) -> Optional[_T]: ...
231
232 @overload
233 def get(self, key: str, default: Union[_DT, _T]) -> Union[_DT, _T]: ...
234
235 def get(
236 self, key: str, default: Optional[Union[_DT, _T]] = None
237 ) -> Optional[Union[_T, _DT]]:
238 if key in self:
239 return self[key]
240 else:
241 return default
242
243 def keys(self) -> List[str]:
244 return list(self._data)
245
246 def values(self) -> List[_T]:
247 return list(self._data.values())
248
249 def items(self) -> List[Tuple[str, _T]]:
250 return list(self._data.items())
251
252 def has_key(self, key: str) -> bool:
253 return key in self._data
254
255 def clear(self) -> None:
256 self._data.clear()
257
258
259class OrderedProperties(Properties[_T]):
260 """Provide a __getattr__/__setattr__ interface with an OrderedDict
261 as backing store."""
262
263 __slots__ = ()
264
265 def __init__(self):
266 Properties.__init__(self, OrderedDict())
267
268
269class ReadOnlyProperties(ReadOnlyContainer, Properties[_T]):
270 """Provide immutable dict/object attribute to an underlying dictionary."""
271
272 __slots__ = ()
273
274
275def _ordered_dictionary_sort(d, key=None):
276 """Sort an OrderedDict in-place."""
277
278 items = [(k, d[k]) for k in sorted(d, key=key)]
279
280 d.clear()
281
282 d.update(items)
283
284
285OrderedDict = dict
286sort_dictionary = _ordered_dictionary_sort
287
288
289class WeakSequence(Sequence[_T]):
290 def __init__(self, __elements: Sequence[_T] = ()):
291 # adapted from weakref.WeakKeyDictionary, prevent reference
292 # cycles in the collection itself
293 def _remove(item, selfref=weakref.ref(self)):
294 self = selfref()
295 if self is not None:
296 self._storage.remove(item)
297
298 self._remove = _remove
299 self._storage = [
300 weakref.ref(element, _remove) for element in __elements
301 ]
302
303 def append(self, item):
304 self._storage.append(weakref.ref(item, self._remove))
305
306 def __len__(self):
307 return len(self._storage)
308
309 def __iter__(self):
310 return (
311 obj for obj in (ref() for ref in self._storage) if obj is not None
312 )
313
314 def __getitem__(self, index):
315 try:
316 obj = self._storage[index]
317 except KeyError:
318 raise IndexError("Index %s out of range" % index)
319 else:
320 return obj()
321
322
323class OrderedIdentitySet(IdentitySet):
324 def __init__(self, iterable: Optional[Iterable[Any]] = None):
325 IdentitySet.__init__(self)
326 self._members = OrderedDict()
327 if iterable:
328 for o in iterable:
329 self.add(o)
330
331
332class PopulateDict(Dict[_KT, _VT]):
333 """A dict which populates missing values via a creation function.
334
335 Note the creation function takes a key, unlike
336 collections.defaultdict.
337
338 """
339
340 def __init__(self, creator: Callable[[_KT], _VT]):
341 self.creator = creator
342
343 def __missing__(self, key: Any) -> Any:
344 self[key] = val = self.creator(key)
345 return val
346
347
348class WeakPopulateDict(Dict[_KT, _VT]):
349 """Like PopulateDict, but assumes a self + a method and does not create
350 a reference cycle.
351
352 """
353
354 def __init__(self, creator_method: types.MethodType):
355 self.creator = creator_method.__func__
356 weakself = creator_method.__self__
357 self.weakself = weakref.ref(weakself)
358
359 def __missing__(self, key: Any) -> Any:
360 self[key] = val = self.creator(self.weakself(), key)
361 return val
362
363
364# Define collections that are capable of storing
365# ColumnElement objects as hashable keys/elements.
366# At this point, these are mostly historical, things
367# used to be more complicated.
368column_set = set
369column_dict = dict
370ordered_column_set = OrderedSet
371
372
373class UniqueAppender(Generic[_T]):
374 """Appends items to a collection ensuring uniqueness.
375
376 Additional appends() of the same object are ignored. Membership is
377 determined by identity (``is a``) not equality (``==``).
378 """
379
380 __slots__ = "data", "_data_appender", "_unique"
381
382 data: Union[Iterable[_T], Set[_T], List[_T]]
383 _data_appender: Callable[[_T], None]
384 _unique: Dict[int, Literal[True]]
385
386 def __init__(
387 self,
388 data: Union[Iterable[_T], Set[_T], List[_T]],
389 via: Optional[str] = None,
390 ):
391 self.data = data
392 self._unique = {}
393 if via:
394 self._data_appender = getattr(data, via)
395 elif hasattr(data, "append"):
396 self._data_appender = cast("List[_T]", data).append
397 elif hasattr(data, "add"):
398 self._data_appender = cast("Set[_T]", data).add
399
400 def append(self, item: _T) -> None:
401 id_ = id(item)
402 if id_ not in self._unique:
403 self._data_appender(item)
404 self._unique[id_] = True
405
406 def __iter__(self) -> Iterator[_T]:
407 return iter(self.data)
408
409
410def coerce_generator_arg(arg: Any) -> List[Any]:
411 if len(arg) == 1 and isinstance(arg[0], types.GeneratorType):
412 return list(arg[0])
413 else:
414 return cast("List[Any]", arg)
415
416
417def to_list(x: Any, default: Optional[List[Any]] = None) -> List[Any]:
418 if x is None:
419 return default # type: ignore
420 if not is_non_string_iterable(x):
421 return [x]
422 elif isinstance(x, list):
423 return x
424 else:
425 return list(x)
426
427
428def has_intersection(set_, iterable):
429 r"""return True if any items of set\_ are present in iterable.
430
431 Goes through special effort to ensure __hash__ is not called
432 on items in iterable that don't support it.
433
434 """
435 # TODO: optimize, write in C, etc.
436 return bool(set_.intersection([i for i in iterable if i.__hash__]))
437
438
439def to_set(x):
440 if x is None:
441 return set()
442 if not isinstance(x, set):
443 return set(to_list(x))
444 else:
445 return x
446
447
448def to_column_set(x: Any) -> Set[Any]:
449 if x is None:
450 return column_set()
451 if not isinstance(x, column_set):
452 return column_set(to_list(x))
453 else:
454 return x
455
456
457def update_copy(d, _new=None, **kw):
458 """Copy the given dict and update with the given values."""
459
460 d = d.copy()
461 if _new:
462 d.update(_new)
463 d.update(**kw)
464 return d
465
466
467def flatten_iterator(x: Iterable[_T]) -> Iterator[_T]:
468 """Given an iterator of which further sub-elements may also be
469 iterators, flatten the sub-elements into a single iterator.
470
471 """
472 elem: _T
473 for elem in x:
474 if not isinstance(elem, str) and hasattr(elem, "__iter__"):
475 yield from flatten_iterator(elem)
476 else:
477 yield elem
478
479
480class LRUCache(typing.MutableMapping[_KT, _VT]):
481 """Dictionary with 'squishy' removal of least
482 recently used items.
483
484 Note that either get() or [] should be used here, but
485 generally its not safe to do an "in" check first as the dictionary
486 can change subsequent to that call.
487
488 """
489
490 __slots__ = (
491 "capacity",
492 "threshold",
493 "size_alert",
494 "_data",
495 "_counter",
496 "_mutex",
497 )
498
499 capacity: int
500 threshold: float
501 size_alert: Optional[Callable[[LRUCache[_KT, _VT]], None]]
502
503 def __init__(
504 self,
505 capacity: int = 100,
506 threshold: float = 0.5,
507 size_alert: Optional[Callable[..., None]] = None,
508 ):
509 self.capacity = capacity
510 self.threshold = threshold
511 self.size_alert = size_alert
512 self._counter = 0
513 self._mutex = threading.Lock()
514 self._data: Dict[_KT, Tuple[_KT, _VT, List[int]]] = {}
515
516 def _inc_counter(self):
517 self._counter += 1
518 return self._counter
519
520 @overload
521 def get(self, key: _KT) -> Optional[_VT]: ...
522
523 @overload
524 def get(self, key: _KT, default: Union[_VT, _T]) -> Union[_VT, _T]: ...
525
526 def get(
527 self, key: _KT, default: Optional[Union[_VT, _T]] = None
528 ) -> Optional[Union[_VT, _T]]:
529 item = self._data.get(key)
530 if item is not None:
531 item[2][0] = self._inc_counter()
532 return item[1]
533 else:
534 return default
535
536 def __getitem__(self, key: _KT) -> _VT:
537 item = self._data[key]
538 item[2][0] = self._inc_counter()
539 return item[1]
540
541 def __iter__(self) -> Iterator[_KT]:
542 return iter(self._data)
543
544 def __len__(self) -> int:
545 return len(self._data)
546
547 def values(self) -> ValuesView[_VT]:
548 return typing.ValuesView({k: i[1] for k, i in self._data.items()})
549
550 def __setitem__(self, key: _KT, value: _VT) -> None:
551 self._data[key] = (key, value, [self._inc_counter()])
552 self._manage_size()
553
554 def __delitem__(self, __v: _KT) -> None:
555 del self._data[__v]
556
557 @property
558 def size_threshold(self) -> float:
559 return self.capacity + self.capacity * self.threshold
560
561 def _manage_size(self) -> None:
562 if not self._mutex.acquire(False):
563 return
564 try:
565 size_alert = bool(self.size_alert)
566 while len(self) > self.capacity + self.capacity * self.threshold:
567 if size_alert:
568 size_alert = False
569 self.size_alert(self) # type: ignore
570 by_counter = sorted(
571 self._data.values(),
572 key=operator.itemgetter(2),
573 reverse=True,
574 )
575 for item in by_counter[self.capacity :]:
576 try:
577 del self._data[item[0]]
578 except KeyError:
579 # deleted elsewhere; skip
580 continue
581 finally:
582 self._mutex.release()
583
584
585class _CreateFuncType(Protocol[_T_co]):
586 def __call__(self) -> _T_co: ...
587
588
589class _ScopeFuncType(Protocol):
590 def __call__(self) -> Any: ...
591
592
593class ScopedRegistry(Generic[_T]):
594 """A Registry that can store one or multiple instances of a single
595 class on the basis of a "scope" function.
596
597 The object implements ``__call__`` as the "getter", so by
598 calling ``myregistry()`` the contained object is returned
599 for the current scope.
600
601 :param createfunc:
602 a callable that returns a new object to be placed in the registry
603
604 :param scopefunc:
605 a callable that will return a key to store/retrieve an object.
606 """
607
608 __slots__ = "createfunc", "scopefunc", "registry"
609
610 createfunc: _CreateFuncType[_T]
611 scopefunc: _ScopeFuncType
612 registry: Any
613
614 def __init__(
615 self, createfunc: Callable[[], _T], scopefunc: Callable[[], Any]
616 ):
617 """Construct a new :class:`.ScopedRegistry`.
618
619 :param createfunc: A creation function that will generate
620 a new value for the current scope, if none is present.
621
622 :param scopefunc: A function that returns a hashable
623 token representing the current scope (such as, current
624 thread identifier).
625
626 """
627 self.createfunc = createfunc
628 self.scopefunc = scopefunc
629 self.registry = {}
630
631 def __call__(self) -> _T:
632 key = self.scopefunc()
633 try:
634 return self.registry[key] # type: ignore[no-any-return]
635 except KeyError:
636 return self.registry.setdefault(key, self.createfunc()) # type: ignore[no-any-return] # noqa: E501
637
638 def has(self) -> bool:
639 """Return True if an object is present in the current scope."""
640
641 return self.scopefunc() in self.registry
642
643 def set(self, obj: _T) -> None:
644 """Set the value for the current scope."""
645
646 self.registry[self.scopefunc()] = obj
647
648 def clear(self) -> None:
649 """Clear the current scope, if any."""
650
651 try:
652 del self.registry[self.scopefunc()]
653 except KeyError:
654 pass
655
656
657class ThreadLocalRegistry(ScopedRegistry[_T]):
658 """A :class:`.ScopedRegistry` that uses a ``threading.local()``
659 variable for storage.
660
661 """
662
663 def __init__(self, createfunc: Callable[[], _T]):
664 self.createfunc = createfunc
665 self.registry = threading.local()
666
667 def __call__(self) -> _T:
668 try:
669 return self.registry.value # type: ignore[no-any-return]
670 except AttributeError:
671 val = self.registry.value = self.createfunc()
672 return val
673
674 def has(self) -> bool:
675 return hasattr(self.registry, "value")
676
677 def set(self, obj: _T) -> None:
678 self.registry.value = obj
679
680 def clear(self) -> None:
681 try:
682 del self.registry.value
683 except AttributeError:
684 pass
685
686
687def has_dupes(sequence, target):
688 """Given a sequence and search object, return True if there's more
689 than one, False if zero or one of them.
690
691
692 """
693 # compare to .index version below, this version introduces less function
694 # overhead and is usually the same speed. At 15000 items (way bigger than
695 # a relationship-bound collection in memory usually is) it begins to
696 # fall behind the other version only by microseconds.
697 c = 0
698 for item in sequence:
699 if item is target:
700 c += 1
701 if c > 1:
702 return True
703 return False
704
705
706# .index version. the two __contains__ calls as well
707# as .index() and isinstance() slow this down.
708# def has_dupes(sequence, target):
709# if target not in sequence:
710# return False
711# elif not isinstance(sequence, collections_abc.Sequence):
712# return False
713#
714# idx = sequence.index(target)
715# return target in sequence[idx + 1:]
716 