codekingpro/portable-devtools
114k
1# cyextension/collections.pyx
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
7cimport cython
8from cpython.long cimport PyLong_FromLongLong
9from cpython.set cimport PySet_Add
10
11from collections.abc import Collection
12from itertools import filterfalse
13
14cdef bint add_not_present(set seen, object item, hashfunc):
15 hash_value = hashfunc(item)
16 if hash_value not in seen:
17 PySet_Add(seen, hash_value)
18 return True
19 else:
20 return False
21
22cdef list cunique_list(seq, hashfunc=None):
23 cdef set seen = set()
24 if not hashfunc:
25 return [x for x in seq if x not in seen and not PySet_Add(seen, x)]
26 else:
27 return [x for x in seq if add_not_present(seen, x, hashfunc)]
28
29def unique_list(seq, hashfunc=None):
30 return cunique_list(seq, hashfunc)
31
32cdef class OrderedSet(set):
33
34 cdef list _list
35
36 @classmethod
37 def __class_getitem__(cls, key):
38 return cls
39
40 def __init__(self, d=None):
41 set.__init__(self)
42 if d is not None:
43 self._list = cunique_list(d)
44 set.update(self, self._list)
45 else:
46 self._list = []
47
48 cpdef OrderedSet copy(self):
49 cdef OrderedSet cp = OrderedSet.__new__(OrderedSet)
50 cp._list = list(self._list)
51 set.update(cp, cp._list)
52 return cp
53
54 @cython.final
55 cdef OrderedSet _from_list(self, list new_list):
56 cdef OrderedSet new = OrderedSet.__new__(OrderedSet)
57 new._list = new_list
58 set.update(new, new_list)
59 return new
60
61 def add(self, element):
62 if element not in self:
63 self._list.append(element)
64 PySet_Add(self, element)
65
66 def remove(self, element):
67 # set.remove will raise if element is not in self
68 set.remove(self, element)
69 self._list.remove(element)
70
71 def pop(self):
72 try:
73 value = self._list.pop()
74 except IndexError:
75 raise KeyError("pop from an empty set") from None
76 set.remove(self, value)
77 return value
78
79 def insert(self, Py_ssize_t pos, element):
80 if element not in self:
81 self._list.insert(pos, element)
82 PySet_Add(self, element)
83
84 def discard(self, element):
85 if element in self:
86 set.remove(self, element)
87 self._list.remove(element)
88
89 def clear(self):
90 set.clear(self)
91 self._list = []
92
93 def __getitem__(self, key):
94 return self._list[key]
95
96 def __iter__(self):
97 return iter(self._list)
98
99 def __add__(self, other):
100 return self.union(other)
101
102 def __repr__(self):
103 return "%s(%r)" % (self.__class__.__name__, self._list)
104
105 __str__ = __repr__
106
107 def update(self, *iterables):
108 for iterable in iterables:
109 for e in iterable:
110 if e not in self:
111 self._list.append(e)
112 set.add(self, e)
113
114 def __ior__(self, iterable):
115 self.update(iterable)
116 return self
117
118 def union(self, *other):
119 result = self.copy()
120 result.update(*other)
121 return result
122
123 def __or__(self, other):
124 return self.union(other)
125
126 def intersection(self, *other):
127 cdef set other_set = set.intersection(self, *other)
128 return self._from_list([a for a in self._list if a in other_set])
129
130 def __and__(self, other):
131 return self.intersection(other)
132
133 def symmetric_difference(self, other):
134 cdef set other_set
135 if isinstance(other, set):
136 other_set = <set> other
137 collection = other_set
138 elif isinstance(other, Collection):
139 collection = other
140 other_set = set(other)
141 else:
142 collection = list(other)
143 other_set = set(collection)
144 result = self._from_list([a for a in self._list if a not in other_set])
145 result.update(a for a in collection if a not in self)
146 return result
147
148 def __xor__(self, other):
149 return self.symmetric_difference(other)
150
151 def difference(self, *other):
152 cdef set other_set = set.difference(self, *other)
153 return self._from_list([a for a in self._list if a in other_set])
154
155 def __sub__(self, other):
156 return self.difference(other)
157
158 def intersection_update(self, *other):
159 set.intersection_update(self, *other)
160 self._list = [a for a in self._list if a in self]
161
162 def __iand__(self, other):
163 self.intersection_update(other)
164 return self
165
166 cpdef symmetric_difference_update(self, other):
167 collection = other if isinstance(other, Collection) else list(other)
168 set.symmetric_difference_update(self, collection)
169 self._list = [a for a in self._list if a in self]
170 self._list += [a for a in collection if a in self]
171
172 def __ixor__(self, other):
173 self.symmetric_difference_update(other)
174 return self
175
176 def difference_update(self, *other):
177 set.difference_update(self, *other)
178 self._list = [a for a in self._list if a in self]
179
180 def __isub__(self, other):
181 self.difference_update(other)
182 return self
183
184cdef object cy_id(object item):
185 return PyLong_FromLongLong(<long long> (<void *>item))
186
187# NOTE: cython 0.x will call __add__, __sub__, etc with the parameter swapped
188# instead of the __rmeth__, so they need to check that also self is of the
189# correct type. This is fixed in cython 3.x. See:
190# https://docs.cython.org/en/latest/src/userguide/special_methods.html#arithmetic-methods
191cdef class IdentitySet:
192 """A set that considers only object id() for uniqueness.
193
194 This strategy has edge cases for builtin types- it's possible to have
195 two 'foo' strings in one of these sets, for example. Use sparingly.
196
197 """
198
199 cdef dict _members
200
201 def __init__(self, iterable=None):
202 self._members = {}
203 if iterable:
204 self.update(iterable)
205
206 def add(self, value):
207 self._members[cy_id(value)] = value
208
209 def __contains__(self, value):
210 return cy_id(value) in self._members
211
212 cpdef remove(self, value):
213 del self._members[cy_id(value)]
214
215 def discard(self, value):
216 try:
217 self.remove(value)
218 except KeyError:
219 pass
220
221 def pop(self):
222 cdef tuple pair
223 try:
224 pair = self._members.popitem()
225 return pair[1]
226 except KeyError:
227 raise KeyError("pop from an empty set")
228
229 def clear(self):
230 self._members.clear()
231
232 def __eq__(self, other):
233 cdef IdentitySet other_
234 if isinstance(other, IdentitySet):
235 other_ = other
236 return self._members == other_._members
237 else:
238 return False
239
240 def __ne__(self, other):
241 cdef IdentitySet other_
242 if isinstance(other, IdentitySet):
243 other_ = other
244 return self._members != other_._members
245 else:
246 return True
247
248 cpdef issubset(self, iterable):
249 cdef IdentitySet other
250 if isinstance(iterable, self.__class__):
251 other = iterable
252 else:
253 other = self.__class__(iterable)
254
255 if len(self) > len(other):
256 return False
257 for m in filterfalse(other._members.__contains__, self._members):
258 return False
259 return True
260
261 def __le__(self, other):
262 if not isinstance(other, IdentitySet):
263 return NotImplemented
264 return self.issubset(other)
265
266 def __lt__(self, other):
267 if not isinstance(other, IdentitySet):
268 return NotImplemented
269 return len(self) < len(other) and self.issubset(other)
270
271 cpdef issuperset(self, iterable):
272 cdef IdentitySet other
273 if isinstance(iterable, self.__class__):
274 other = iterable
275 else:
276 other = self.__class__(iterable)
277
278 if len(self) < len(other):
279 return False
280 for m in filterfalse(self._members.__contains__, other._members):
281 return False
282 return True
283
284 def __ge__(self, other):
285 if not isinstance(other, IdentitySet):
286 return NotImplemented
287 return self.issuperset(other)
288
289 def __gt__(self, other):
290 if not isinstance(other, IdentitySet):
291 return NotImplemented
292 return len(self) > len(other) and self.issuperset(other)
293
294 cpdef IdentitySet union(self, iterable):
295 cdef IdentitySet result = self.__class__()
296 result._members.update(self._members)
297 result.update(iterable)
298 return result
299
300 def __or__(self, other):
301 if not isinstance(other, IdentitySet) or not isinstance(self, IdentitySet):
302 return NotImplemented
303 return self.union(other)
304
305 cpdef update(self, iterable):
306 for obj in iterable:
307 self._members[cy_id(obj)] = obj
308
309 def __ior__(self, other):
310 if not isinstance(other, IdentitySet):
311 return NotImplemented
312 self.update(other)
313 return self
314
315 cpdef IdentitySet difference(self, iterable):
316 cdef IdentitySet result = self.__new__(self.__class__)
317 if isinstance(iterable, self.__class__):
318 other = (<IdentitySet>iterable)._members
319 else:
320 other = {cy_id(obj) for obj in iterable}
321 result._members = {k:v for k, v in self._members.items() if k not in other}
322 return result
323
324 def __sub__(self, other):
325 if not isinstance(other, IdentitySet) or not isinstance(self, IdentitySet):
326 return NotImplemented
327 return self.difference(other)
328
329 cpdef difference_update(self, iterable):
330 cdef IdentitySet other = self.difference(iterable)
331 self._members = other._members
332
333 def __isub__(self, other):
334 if not isinstance(other, IdentitySet):
335 return NotImplemented
336 self.difference_update(other)
337 return self
338
339 cpdef IdentitySet intersection(self, iterable):
340 cdef IdentitySet result = self.__new__(self.__class__)
341 if isinstance(iterable, self.__class__):
342 other = (<IdentitySet>iterable)._members
343 else:
344 other = {cy_id(obj) for obj in iterable}
345 result._members = {k: v for k, v in self._members.items() if k in other}
346 return result
347
348 def __and__(self, other):
349 if not isinstance(other, IdentitySet) or not isinstance(self, IdentitySet):
350 return NotImplemented
351 return self.intersection(other)
352
353 cpdef intersection_update(self, iterable):
354 cdef IdentitySet other = self.intersection(iterable)
355 self._members = other._members
356
357 def __iand__(self, other):
358 if not isinstance(other, IdentitySet):
359 return NotImplemented
360 self.intersection_update(other)
361 return self
362
363 cpdef IdentitySet symmetric_difference(self, iterable):
364 cdef IdentitySet result = self.__new__(self.__class__)
365 cdef dict other
366 if isinstance(iterable, self.__class__):
367 other = (<IdentitySet>iterable)._members
368 else:
369 other = {cy_id(obj): obj for obj in iterable}
370 result._members = {k: v for k, v in self._members.items() if k not in other}
371 result._members.update(
372 [(k, v) for k, v in other.items() if k not in self._members]
373 )
374 return result
375
376 def __xor__(self, other):
377 if not isinstance(other, IdentitySet) or not isinstance(self, IdentitySet):
378 return NotImplemented
379 return self.symmetric_difference(other)
380
381 cpdef symmetric_difference_update(self, iterable):
382 cdef IdentitySet other = self.symmetric_difference(iterable)
383 self._members = other._members
384
385 def __ixor__(self, other):
386 if not isinstance(other, IdentitySet):
387 return NotImplemented
388 self.symmetric_difference(other)
389 return self
390
391 cpdef IdentitySet copy(self):
392 cdef IdentitySet cp = self.__new__(self.__class__)
393 cp._members = self._members.copy()
394 return cp
395
396 def __copy__(self):
397 return self.copy()
398
399 def __len__(self):
400 return len(self._members)
401
402 def __iter__(self):
403 return iter(self._members.values())
404
405 def __hash__(self):
406 raise TypeError("set objects are unhashable")
407
408 def __repr__(self):
409 return "%s(%r)" % (type(self).__name__, list(self._members.values()))
410 