Team Ai
Datasetpublic

codekingpro/portable-devtools

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