1use super::{
5 IterStatus, PositionIterInternal, PyDict, PyDictRef, PyGenericAlias, PyTupleRef, PyType,
6 PyTypeRef, builtins_iter, locked_step,
7};
8use crate::{
9 AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, TryFromObject,
10 atomic_func,
11 class::{PyClassDef, PyClassImpl},
12 common::{
13 ascii,
14 hash::PyHash,
15 lock::{LazyLock, PyMutex},
16 rc::PyRc,
17 wtf8::Wtf8Buf,
18 },
19 convert::ToPyResult,
20 dict_inner::{self, DictSize},
21 function::{
22 ArgIterable, FuncArgs, NameOthers, OptionalArg, PosArgs, PyArithmeticValue,
23 PyComparisonValue,
24 },
25 protocol::{PyIterReturn, PyNumberMethods, PySequenceMethods},
26 recursion::ReprGuard,
27 types::AsNumber,
28 types::{
29 AsSequence, Comparable, Constructor, DefaultConstructor, Hashable, Initializer, IterNext,
30 Iterable, PyComparisonOp, Representable, SelfIter,
31 },
32 utils::collection_repr,
33 vm::VirtualMachine,
34};
35use core::{borrow::Borrow, fmt};
36use rustpython_common::{
37 atomic::{Ordering, PyAtomic, Radium},
38 hash,
39};
40
41pub(crate) type SetContentType = dict_inner::Dict<()>;
42
43#[pyclass(module = false, name = "set", unhashable = true, traverse)]
44#[derive(Default)]
45pub struct PySet {
46 pub(super) inner: PySetInner,
47}
48
49impl PySet {
50 #[deprecated(note = "Use `PySet::default().into_ref(ctx)` instead")]
51 pub fn new_ref(ctx: &Context) -> PyRef<Self> {
52 Self::default().into_ref(ctx)
53 }
54
55 #[must_use]
56 pub fn elements(&self) -> Vec<PyObjectRef> {
57 self.inner.elements()
58 }
59
60 fn fold_op(
61 &self,
62 others: impl core::iter::Iterator<Item = ArgIterable>,
63 op: fn(&PySetInner, ArgIterable, &VirtualMachine) -> PyResult<PySetInner>,
64 vm: &VirtualMachine,
65 ) -> PyResult<Self> {
66 Ok(Self {
67 inner: self.inner.fold_op(others, op, vm)?,
68 })
69 }
70
71 fn op(
72 &self,
73 other: AnySet,
74 op: fn(&PySetInner, ArgIterable, &VirtualMachine) -> PyResult<PySetInner>,
75 vm: &VirtualMachine,
76 ) -> PyResult<Self> {
77 Ok(Self {
78 inner: self
79 .inner
80 .fold_op(core::iter::once(other.into_iterable(vm)?), op, vm)?,
81 })
82 }
83}
84
85#[pyclass(module = false, name = "frozenset", unhashable = true)]
86pub struct PyFrozenSet {
87 inner: PySetInner,
88 hash: PyAtomic<PyHash>,
89}
90
91impl Default for PyFrozenSet {
92 fn default() -> Self {
93 Self {
94 inner: PySetInner::default(),
95 hash: hash::SENTINEL.into(),
96 }
97 }
98}
99
100impl PyFrozenSet {
101 pub fn from_iter(
103 vm: &VirtualMachine,
104 it: impl IntoIterator<Item = PyObjectRef>,
105 ) -> PyResult<Self> {
106 let inner = PySetInner::default();
107 for elem in it {
108 inner.add(&elem, vm)?;
109 }
110 Ok(Self {
112 inner,
113 ..Default::default()
114 })
115 }
116
117 pub fn elements(&self) -> Vec<PyObjectRef> {
118 self.inner.elements()
119 }
120
121 fn fold_op(
122 &self,
123 others: impl core::iter::Iterator<Item = ArgIterable>,
124 op: fn(&PySetInner, ArgIterable, &VirtualMachine) -> PyResult<PySetInner>,
125 vm: &VirtualMachine,
126 ) -> PyResult<Self> {
127 Ok(Self {
128 inner: self.inner.fold_op(others, op, vm)?,
129 ..Default::default()
130 })
131 }
132
133 fn op(
134 &self,
135 other: AnySet,
136 op: fn(&PySetInner, ArgIterable, &VirtualMachine) -> PyResult<PySetInner>,
137 vm: &VirtualMachine,
138 ) -> PyResult<Self> {
139 Ok(Self {
140 inner: self
141 .inner
142 .fold_op(core::iter::once(other.into_iterable(vm)?), op, vm)?,
143 ..Default::default()
144 })
145 }
146}
147
148impl fmt::Debug for PySet {
149 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
150 f.write_str("set")
152 }
153}
154
155impl fmt::Debug for PyFrozenSet {
156 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
157 f.write_str("PyFrozenSet ")?;
159 f.debug_set().entries(self.elements().iter()).finish()
160 }
161}
162
163impl PyPayload for PySet {
164 #[inline]
165 fn class(ctx: &Context) -> &'static Py<PyType> {
166 ctx.types.set_type
167 }
168}
169
170impl PyPayload for PyFrozenSet {
171 #[inline]
172 fn class(ctx: &Context) -> &'static Py<PyType> {
173 ctx.types.frozenset_type
174 }
175}
176
177#[derive(Default, Clone)]
178pub(super) struct PySetInner {
179 content: PyRc<SetContentType>,
180}
181
182unsafe impl crate::object::Traverse for PySetInner {
183 fn traverse(&self, tracer_fn: &mut crate::object::TraverseFn<'_>) {
184 self.content.traverse(tracer_fn)
186 }
187}
188
189impl PySetInner {
190 pub(super) fn from_iter<T>(iter: T, vm: &VirtualMachine) -> PyResult<Self>
191 where
192 T: IntoIterator<Item = PyResult<PyObjectRef>>,
193 {
194 let set = Self::default();
195 for item in iter {
196 let item = item?;
197 set.add(&item, vm)?;
198 }
199 Ok(set)
200 }
201
202 fn from_object(iterable: PyObjectRef, vm: &VirtualMachine) -> PyResult<Self> {
205 let set = Self::default();
206 set.update_internal(iterable, vm)?;
207 Ok(set)
208 }
209
210 fn cached_hashes(obj: &PyObject, vm: &VirtualMachine) -> Option<Vec<(PyObjectRef, PyHash)>> {
214 if let Some(set) = extract_set(obj) {
215 Some(set.content.keys_with_hashes())
216 } else {
217 obj.downcast_ref_if_exact::<PyDict>(vm)
218 .map(|dict| dict._as_dict_inner().keys_with_hashes())
219 }
220 }
221
222 fn fold_op<O>(
223 &self,
224 others: impl core::iter::Iterator<Item = O>,
225 op: fn(&Self, O, &VirtualMachine) -> PyResult<Self>,
226 vm: &VirtualMachine,
227 ) -> PyResult<Self> {
228 let mut res = self.copy();
229 for other in others {
230 res = op(&res, other, vm)?;
231 }
232 Ok(res)
233 }
234
235 fn intersection_multi(
236 &self,
237 mut others: impl core::iter::Iterator<Item = ArgIterable>,
238 vm: &VirtualMachine,
239 ) -> PyResult<Self> {
240 let Some(other) = others.next() else {
241 return Ok(self.copy());
242 };
243 let mut result = self.intersection(other, vm)?;
244 for other in others {
245 result = result.intersection(other, vm)?;
246 }
247 Ok(result)
248 }
249
250 fn difference_multi(
251 &self,
252 mut others: impl core::iter::Iterator<Item = ArgIterable>,
253 vm: &VirtualMachine,
254 ) -> PyResult<Self> {
255 let Some(other) = others.next() else {
256 return Ok(self.copy());
257 };
258 let result = self.difference_new(other, vm)?;
259 result.difference_update(others, vm)?;
260 Ok(result)
261 }
262
263 fn len(&self) -> usize {
264 self.content.len()
265 }
266
267 fn sizeof(&self) -> usize {
268 self.content.sizeof()
269 }
270
271 fn copy(&self) -> Self {
272 Self {
273 content: PyRc::new((*self.content).clone()),
274 }
275 }
276
277 fn contains(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
278 let result = self
279 .retry_op_with_frozenset(needle, vm, |needle, vm| self.content.contains(vm, needle));
280 Self::wrap_unhashable_error(result, needle, vm)
281 }
282
283 fn contains_known_hash(
285 &self,
286 needle: &PyObject,
287 hash: PyHash,
288 vm: &VirtualMachine,
289 ) -> PyResult<bool> {
290 self.content.contains_known_hash(vm, needle, hash)
291 }
292
293 fn compare(&self, other: &Self, op: PyComparisonOp, vm: &VirtualMachine) -> PyResult<bool> {
294 if op == PyComparisonOp::Ne {
295 return self.compare(other, PyComparisonOp::Eq, vm).map(|eq| !eq);
296 }
297 if !op.eval_ord(self.len().cmp(&other.len())) {
298 return Ok(false);
299 }
300
301 let (superset, subset) = match op {
302 PyComparisonOp::Lt | PyComparisonOp::Le | PyComparisonOp::Eq => (other, self),
303 _ => (self, other),
304 };
305
306 for (key, hash) in subset.content.keys_with_hashes() {
307 if !superset.contains_known_hash(&key, hash, vm)? {
308 return Ok(false);
309 }
310 }
311 Ok(true)
312 }
313
314 pub(super) fn union(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<Self> {
315 let set = self.clone();
316 if let Some(elements) = Self::cached_hashes(other.as_object(), vm) {
317 for (item, hash) in elements {
318 set.add_known_hash(&item, hash, vm)?;
319 }
320 return Ok(set);
321 }
322 for item in other.iter(vm)? {
323 let item = item?;
324 set.add(&item, vm)?;
325 }
326
327 Ok(set)
328 }
329
330 pub(super) fn intersection(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<Self> {
331 if let Some(other_set) = extract_set(other.as_object()) {
332 return self.intersection_set(other_set, vm);
333 }
334 let set = Self::default();
335 for item in other.iter(vm)? {
336 let obj = item?;
337 let hash = obj.hash(vm)?;
338 if self.contains_known_hash(&obj, hash, vm)? {
339 set.add_known_hash(&obj, hash, vm)?;
340 if set.len() >= self.len() {
341 break;
342 }
343 }
344 }
345 Ok(set)
346 }
347
348 fn intersection_set(&self, other: &Self, vm: &VirtualMachine) -> PyResult<Self> {
349 if PyRc::ptr_eq(&self.content, &other.content) {
350 return Ok(self.copy());
351 }
352 let (target, source) = if self.len() < other.len() {
353 (other, self)
354 } else {
355 (self, other)
356 };
357 let set = Self::default();
358 for (obj, hash) in source.content.keys_with_hashes() {
359 if target.contains_known_hash(&obj, hash, vm)? {
360 set.add_known_hash(&obj, hash, vm)?;
361 }
362 }
363 Ok(set)
364 }
365
366 fn difference_new(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<Self> {
367 if let Some(other_set) = extract_set(other.as_object()) {
370 if self.len() >> 2 <= other_set.len() {
371 return self
372 .difference_by(|key, hash| other_set.contains_known_hash(key, hash, vm), vm);
373 }
374 } else if let Some(dict) = other.as_object().downcast_ref_if_exact::<PyDict>(vm)
375 && self.len() >> 2 <= dict._as_dict_inner().len()
376 {
377 return self.difference_by(
378 |key, hash| dict._as_dict_inner().contains_known_hash(vm, key, hash),
379 vm,
380 );
381 }
382 self.copy().difference(other, vm)
383 }
384
385 pub(super) fn difference(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<Self> {
387 self.difference_update(core::iter::once(other), vm)?;
388 Ok(self.clone())
389 }
390
391 fn difference_by(
392 &self,
393 contains: impl Fn(&PyObject, PyHash) -> PyResult<bool>,
394 vm: &VirtualMachine,
395 ) -> PyResult<Self> {
396 let result = Self::default();
397 for (key, hash) in self.content.keys_with_hashes() {
398 if !contains(&key, hash)? {
399 result.add_known_hash(&key, hash, vm)?;
400 }
401 }
402 Ok(result)
403 }
404
405 pub(super) fn symmetric_difference(
406 &self,
407 other: ArgIterable,
408 vm: &VirtualMachine,
409 ) -> PyResult<Self> {
410 let new_inner = self.clone();
411
412 if let Some(elements) = Self::cached_hashes(other.as_object(), vm) {
413 for (item, hash) in elements {
415 new_inner
416 .content
417 .delete_or_insert_known_hash(vm, &item, hash, ())?;
418 }
419 return Ok(new_inner);
420 }
421
422 let other_set = Self::from_iter(other.iter(vm)?, vm)?;
424
425 for (item, hash) in other_set.content.keys_with_hashes() {
426 new_inner
427 .content
428 .delete_or_insert_known_hash(vm, &item, hash, ())?;
429 }
430
431 Ok(new_inner)
432 }
433
434 fn issuperset(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
435 if let Some(other_set) = extract_set(other.as_object()) {
436 return self.compare(other_set, PyComparisonOp::Ge, vm);
437 }
438 for item in other.iter(vm)? {
439 if !self.contains(&*item?, vm)? {
440 return Ok(false);
441 }
442 }
443 Ok(true)
444 }
445
446 fn issubset(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
447 if let Some(other_set) = extract_set(other.as_object()) {
448 return self.compare(other_set, PyComparisonOp::Le, vm);
449 }
450 Ok(self.intersection(other, vm)?.len() == self.len())
451 }
452
453 pub(super) fn isdisjoint(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
454 if let Some(other_set) = extract_set(other.as_object()) {
455 if core::ptr::eq(self, other_set) {
456 return Ok(self.len() == 0);
457 }
458 let other_type = other.as_object().class();
459 if other_type.is(vm.ctx.types.set_type) || other_type.is(vm.ctx.types.frozenset_type) {
460 let (target, source) = if self.len() < other_set.len() {
461 (other_set, self)
462 } else {
463 (self, other_set)
464 };
465 for (key, hash) in source.content.keys_with_hashes() {
466 if target.contains_known_hash(&key, hash, vm)? {
467 return Ok(false);
468 }
469 }
470 return Ok(true);
471 }
472 }
473 for item in other.iter(vm)? {
474 if self.contains(&*item?, vm)? {
475 return Ok(false);
476 }
477 }
478 Ok(true)
479 }
480
481 fn repr(&self, class_name: Option<&str>, vm: &VirtualMachine) -> PyResult<Wtf8Buf> {
482 let empty = format!("{}()", class_name.unwrap_or("set"));
483 collection_repr(
484 class_name,
485 "{",
486 "}",
487 &empty,
488 self.elements().iter().map(|o| &**o),
489 vm,
490 )
491 }
492
493 fn add(&self, item: &PyObject, vm: &VirtualMachine) -> PyResult<()> {
494 let result = self.content.insert(vm, item, ());
495 Self::wrap_unhashable_error(result, item, vm)
496 }
497
498 fn add_known_hash(&self, item: &PyObject, hash: PyHash, vm: &VirtualMachine) -> PyResult<()> {
500 let result = self.content.insert_known_hash(vm, item, hash, ());
501 Self::wrap_unhashable_error(result, item, vm)
502 }
503
504 fn remove(&self, item: &PyObject, vm: &VirtualMachine) -> PyResult<()> {
505 let result =
506 self.retry_op_with_frozenset(item, vm, |item, vm| self.content.delete(vm, item));
507 Self::wrap_unhashable_error(result, item, vm)
508 }
509
510 fn discard(&self, item: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
511 let result = self
512 .retry_op_with_frozenset(item, vm, |item, vm| self.content.delete_if_exists(vm, item));
513 Self::wrap_unhashable_error(result, item, vm)
514 }
515
516 fn clear(&self) {
517 self.content.clear()
518 }
519
520 fn elements(&self) -> Vec<PyObjectRef> {
521 self.content.keys()
522 }
523
524 fn pop(&self, vm: &VirtualMachine) -> PyResult {
525 if let Some((key, _)) = self.content.pop_back() {
527 Ok(key)
528 } else {
529 let err_msg = vm.ctx.new_str(ascii!("pop from an empty set")).into();
530 Err(vm.new_key_error(err_msg))
531 }
532 }
533
534 fn update_internal(&self, iterable: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
535 if let Ok(any_set) = AnySet::try_from_object(vm, iterable.to_owned()) {
537 self.merge_set(any_set, vm)
538 } else if let Ok(dict) = iterable.to_owned().downcast_exact::<PyDict>(vm) {
540 self.merge_dict(&dict, vm)
541 } else {
542 for item in iterable.try_into_value::<ArgIterable>(vm)?.iter(vm)? {
544 let item = item?;
545 self.add(&item, vm)?;
546 }
547 Ok(())
548 }
549 }
550
551 fn merge_set(&self, any_set: AnySet, vm: &VirtualMachine) -> PyResult<()> {
552 for (item, hash) in any_set.as_inner().content.keys_with_hashes() {
553 self.add_known_hash(&item, hash, vm)?;
554 }
555 Ok(())
556 }
557
558 fn merge_dict(&self, dict: &Py<PyDict>, vm: &VirtualMachine) -> PyResult<()> {
559 for (key, hash) in dict._as_dict_inner().keys_with_hashes() {
560 self.add_known_hash(&key, hash, vm)?;
561 }
562 Ok(())
563 }
564
565 fn intersection_update(
566 &self,
567 others: impl core::iter::Iterator<Item = ArgIterable>,
568 vm: &VirtualMachine,
569 ) -> PyResult<()> {
570 let temp_inner = self.intersection_multi(others, vm)?;
571 let content = PyRc::try_unwrap(temp_inner.content).unwrap_or_else(|table| (*table).clone());
572 self.content.replace_contents(content);
573 Ok(())
574 }
575
576 fn difference_update(
577 &self,
578 others: impl core::iter::Iterator<Item = ArgIterable>,
579 vm: &VirtualMachine,
580 ) -> PyResult<()> {
581 for iterable in others {
582 let elements = if let Some(other_set) = extract_set(iterable.as_object()) {
583 if PyRc::ptr_eq(&self.content, &other_set.content) {
584 self.clear();
585 continue;
586 }
587 Some(if other_set.len() >> 3 > self.len() {
590 self.intersection_set(other_set, vm)?
591 .content
592 .keys_with_hashes()
593 } else {
594 other_set.content.keys_with_hashes()
595 })
596 } else {
597 Self::cached_hashes(iterable.as_object(), vm)
598 };
599 if let Some(elements) = elements {
600 for (item, hash) in elements {
601 self.content.delete_if_exists_known_hash(vm, &*item, hash)?;
602 }
603 continue;
604 }
605 for item in iterable.iter(vm)? {
606 self.content.delete_if_exists(vm, &*item?)?;
607 }
608 }
609 Ok(())
610 }
611
612 fn symmetric_difference_update(
613 &self,
614 others: impl core::iter::Iterator<Item = ArgIterable>,
615 vm: &VirtualMachine,
616 ) -> PyResult<()> {
617 for iterable in others {
618 if let Some(elements) = Self::cached_hashes(iterable.as_object(), vm) {
619 for (item, hash) in elements {
621 self.content
622 .delete_or_insert_known_hash(vm, &item, hash, ())?;
623 }
624 continue;
625 }
626 let iterable_set = Self::from_iter(iterable.iter(vm)?, vm)?;
628 for (item, hash) in iterable_set.content.keys_with_hashes() {
629 self.content
630 .delete_or_insert_known_hash(vm, &item, hash, ())?;
631 }
632 }
633 Ok(())
634 }
635
636 fn hash(&self) -> PyHash {
637 let hasher = self.content.fold_hashes(
638 hash::FrozenSetHash::new(self.len()),
639 |mut hasher, element_hash| {
640 hasher.add(element_hash);
641 hasher
642 },
643 );
644 hasher.finish()
645 }
646
647 fn retry_op_with_frozenset<T, F>(
651 &self,
652 item: &PyObject,
653 vm: &VirtualMachine,
654 op: F,
655 ) -> PyResult<T>
656 where
657 F: Fn(&PyObject, &VirtualMachine) -> PyResult<T>,
658 {
659 op(item, vm).or_else(|original_err| {
660 item.downcast_ref::<PySet>()
661 .ok_or(original_err)
663 .and_then(|set| {
664 op(
665 &PyFrozenSet {
666 inner: set.inner.copy(),
667 ..Default::default()
668 }
669 .into_pyobject(vm),
670 vm,
671 )
672 .map_err(|op_err| {
674 if op_err.fast_isinstance(vm.ctx.exceptions.key_error) {
675 vm.new_key_error(item.to_owned())
676 } else {
677 op_err
678 }
679 })
680 })
681 })
682 }
683
684 fn wrap_unhashable_error<T>(
685 result: PyResult<T>,
686 item: &PyObject,
687 vm: &VirtualMachine,
688 ) -> PyResult<T> {
689 match result {
690 Err(cause) if cause.fast_isinstance(vm.ctx.exceptions.type_error) => {
691 let message = cause.as_object().str(vm)?;
692 let err = vm.new_type_error(format!(
693 "cannot use '{}' as a set element ({message})",
694 item.class().name()
695 ));
696 err.set_cause(Some(cause));
697 Err(err)
698 }
699 result => result,
700 }
701 }
702}
703
704fn extract_set(obj: &PyObject) -> Option<&PySetInner> {
705 match_class!(match obj {
706 ref set @ PySet => Some(&set.inner),
707 ref frozen @ PyFrozenSet => Some(&frozen.inner),
708 _ => None,
709 })
710}
711
712pub(super) fn exact_set_keys_with_hashes(
716 obj: &PyObject,
717 vm: &VirtualMachine,
718) -> Option<Vec<(PyObjectRef, PyHash)>> {
719 let inner = obj
720 .downcast_ref_if_exact::<PySet>(vm)
721 .map(|set| &set.inner)
722 .or_else(|| {
723 obj.downcast_ref_if_exact::<PyFrozenSet>(vm)
724 .map(|frozen| &frozen.inner)
725 })?;
726 Some(inner.content.keys_with_hashes())
727}
728
729fn reduce_set(zelf: &PyObject, vm: &VirtualMachine) -> (PyTypeRef, PyTupleRef, Option<PyDictRef>) {
730 (
731 zelf.class().to_owned(),
732 #[expect(clippy::or_fun_call, reason = "changing this won't compile")]
733 vm.new_tuple((extract_set(zelf)
734 .unwrap_or(&PySetInner::default())
735 .elements(),)),
736 zelf.dict(),
737 )
738}
739
740impl PySet {
741 fn __len__(&self) -> usize {
742 self.inner.len()
743 }
744
745 pub fn contains(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
746 self.inner.contains(needle, vm)
747 }
748
749 fn __or__(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyArithmeticValue<Self>> {
750 if let Ok(other) = AnySet::try_from_object(vm, other) {
751 Ok(PyArithmeticValue::Implemented(self.op(
752 other,
753 PySetInner::union,
754 vm,
755 )?))
756 } else {
757 Ok(PyArithmeticValue::NotImplemented)
758 }
759 }
760
761 fn __and__(
762 &self,
763 other: PyObjectRef,
764 vm: &VirtualMachine,
765 ) -> PyResult<PyArithmeticValue<Self>> {
766 if let Ok(other) = AnySet::try_from_object(vm, other) {
767 Ok(PyArithmeticValue::Implemented(Self {
768 inner: self.inner.intersection(other.into_iterable(vm)?, vm)?,
769 }))
770 } else {
771 Ok(PyArithmeticValue::NotImplemented)
772 }
773 }
774
775 fn __sub__(
776 &self,
777 other: PyObjectRef,
778 vm: &VirtualMachine,
779 ) -> PyResult<PyArithmeticValue<Self>> {
780 if let Ok(other) = AnySet::try_from_object(vm, other) {
781 Ok(PyArithmeticValue::Implemented(Self {
782 inner: self.inner.difference_new(other.into_iterable(vm)?, vm)?,
783 }))
784 } else {
785 Ok(PyArithmeticValue::NotImplemented)
786 }
787 }
788
789 fn __rsub__(
790 zelf: PyRef<Self>,
791 other: PyObjectRef,
792 vm: &VirtualMachine,
793 ) -> PyResult<PyArithmeticValue<Self>> {
794 if let Ok(other) = AnySet::try_from_object(vm, other) {
795 Ok(PyArithmeticValue::Implemented(Self {
796 inner: other
797 .as_inner()
798 .difference_new(ArgIterable::try_from_object(vm, zelf.into())?, vm)?,
799 }))
800 } else {
801 Ok(PyArithmeticValue::NotImplemented)
802 }
803 }
804
805 fn __xor__(
806 &self,
807 other: PyObjectRef,
808 vm: &VirtualMachine,
809 ) -> PyResult<PyArithmeticValue<Self>> {
810 if let Ok(other) = AnySet::try_from_object(vm, other) {
811 Ok(PyArithmeticValue::Implemented(self.op(
812 other,
813 PySetInner::symmetric_difference,
814 vm,
815 )?))
816 } else {
817 Ok(PyArithmeticValue::NotImplemented)
818 }
819 }
820
821 fn __ior__(zelf: PyRef<Self>, set: AnySet, vm: &VirtualMachine) -> PyResult<PyRef<Self>> {
822 zelf.inner.merge_set(set, vm)?;
823 Ok(zelf)
824 }
825
826 fn __iand__(zelf: PyRef<Self>, set: AnySet, vm: &VirtualMachine) -> PyResult<PyRef<Self>> {
827 if !set.is(zelf.as_object()) {
828 zelf.inner
829 .intersection_update(core::iter::once(set.into_iterable(vm)?), vm)?;
830 }
831 Ok(zelf)
832 }
833
834 fn __isub__(zelf: PyRef<Self>, set: AnySet, vm: &VirtualMachine) -> PyResult<PyRef<Self>> {
835 if set.is(zelf.as_object()) {
836 zelf.inner.clear();
837 } else {
838 zelf.inner
839 .difference_update(set.into_iterable_iter(vm)?, vm)?;
840 }
841 Ok(zelf)
842 }
843
844 fn __ixor__(zelf: PyRef<Self>, set: AnySet, vm: &VirtualMachine) -> PyResult<PyRef<Self>> {
845 if set.is(zelf.as_object()) {
846 zelf.inner.clear();
847 } else {
848 zelf.inner
849 .symmetric_difference_update(set.into_iterable_iter(vm)?, vm)?;
850 }
851 Ok(zelf)
852 }
853
854 pub fn add(&self, object: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
855 self.inner.add(&object, vm)
856 }
857}
858
859#[pyclass(
860 with(
861 Constructor,
862 Initializer,
863 AsSequence,
864 Comparable,
865 Iterable,
866 AsNumber,
867 Representable
868 ),
869 flags(BASETYPE, _MATCH_SELF, HAS_WEAKREF)
870)]
871impl Py<PySet> {
872 #[pymethod(coexist)]
873 fn __contains__(&self, object: PyObjectRef, vm: &VirtualMachine) -> PyResult<bool> {
874 self.contains(&object, vm)
875 }
876
877 #[pymethod]
878 fn __sizeof__(&self) -> usize {
879 core::mem::size_of::<PySet>() + self.inner.sizeof()
880 }
881
882 #[pymethod]
883 fn copy(&self) -> PySet {
884 PySet {
885 inner: self.inner.copy(),
886 }
887 }
888
889 #[pymethod]
890 fn union(
891 &self,
892 others: PosArgs<ArgIterable, NameOthers>,
893 vm: &VirtualMachine,
894 ) -> PyResult<PySet> {
895 self.fold_op(others.into_iter(), PySetInner::union, vm)
896 }
897
898 #[pymethod]
899 fn intersection(
900 &self,
901 others: PosArgs<ArgIterable, NameOthers>,
902 vm: &VirtualMachine,
903 ) -> PyResult<PySet> {
904 Ok(PySet {
905 inner: self.inner.intersection_multi(others.into_iter(), vm)?,
906 })
907 }
908
909 #[pymethod]
910 fn difference(
911 &self,
912 others: PosArgs<ArgIterable, NameOthers>,
913 vm: &VirtualMachine,
914 ) -> PyResult<PySet> {
915 Ok(PySet {
916 inner: self.inner.difference_multi(others.into_iter(), vm)?,
917 })
918 }
919
920 #[pymethod]
921 fn symmetric_difference(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<PySet> {
922 self.fold_op(
923 core::iter::once(other),
924 PySetInner::symmetric_difference,
925 vm,
926 )
927 }
928
929 #[pymethod]
930 fn issubset(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
931 self.inner.issubset(other, vm)
932 }
933
934 #[pymethod]
935 fn issuperset(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
936 self.inner.issuperset(other, vm)
937 }
938
939 #[pymethod]
940 fn isdisjoint(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
941 self.inner.isdisjoint(other, vm)
942 }
943
944 #[pymethod]
945 pub fn add(&self, object: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
946 self.payload.add(object, vm)
947 }
948
949 #[pymethod]
950 fn remove(&self, object: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
951 self.inner.remove(&object, vm)
952 }
953
954 #[pymethod]
955 pub fn discard(&self, object: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
956 self.inner.discard(&object, vm).map(|_| ())
957 }
958
959 #[pymethod]
960 pub fn clear(&self) {
961 self.inner.clear()
962 }
963
964 #[pymethod]
965 pub fn pop(&self, vm: &VirtualMachine) -> PyResult {
966 self.inner.pop(vm)
967 }
968
969 #[pymethod]
970 fn update(
971 &self,
972 others: PosArgs<PyObjectRef, NameOthers>,
973 vm: &VirtualMachine,
974 ) -> PyResult<()> {
975 for iterable in others {
976 self.inner.update_internal(iterable, vm)?;
977 }
978 Ok(())
979 }
980
981 #[pymethod]
982 fn intersection_update(
983 &self,
984 others: PosArgs<ArgIterable, NameOthers>,
985 vm: &VirtualMachine,
986 ) -> PyResult<()> {
987 self.inner.intersection_update(others.into_iter(), vm)?;
988 Ok(())
989 }
990
991 #[pymethod]
992 fn difference_update(
993 &self,
994 others: PosArgs<ArgIterable, NameOthers>,
995 vm: &VirtualMachine,
996 ) -> PyResult<()> {
997 self.inner.difference_update(others.into_iter(), vm)
998 }
999
1000 #[pymethod]
1001 fn symmetric_difference_update(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<()> {
1002 self.inner
1003 .symmetric_difference_update(core::iter::once(other), vm)
1004 }
1005
1006 #[pymethod]
1007 fn __reduce__(
1008 zelf: PyRef<PySet>,
1009 vm: &VirtualMachine,
1010 ) -> (PyTypeRef, PyTupleRef, Option<PyDictRef>) {
1011 reduce_set(zelf.as_ref(), vm)
1012 }
1013
1014 #[pyclassmethod]
1015 fn __class_getitem__(
1016 cls: PyTypeRef,
1017 object: PyObjectRef,
1018 vm: &VirtualMachine,
1019 ) -> PyResult<PyGenericAlias> {
1020 PyGenericAlias::from_args(cls, object, vm)
1021 }
1022}
1023
1024impl DefaultConstructor for PySet {}
1025
1026impl Initializer for PySet {
1027 type Args = crate::function::PositionalIterable;
1028
1029 fn init(zelf: &Py<Self>, args: Self::Args, vm: &VirtualMachine) -> PyResult<()> {
1030 zelf.clear();
1031 if let OptionalArg::Present(it) = args.iterable {
1032 zelf.update(PosArgs::<PyObjectRef, NameOthers>::named(vec![it]), vm)?;
1033 }
1034 Ok(())
1035 }
1036}
1037
1038impl AsSequence for PySet {
1039 fn as_sequence() -> &'static PySequenceMethods {
1040 static AS_SEQUENCE: LazyLock<PySequenceMethods> = LazyLock::new(|| PySequenceMethods {
1041 length: atomic_func!(|seq, _vm| Ok(PySet::sequence_downcast(seq).__len__())),
1042 contains: atomic_func!(
1043 |seq, needle, vm| PySet::sequence_downcast(seq).contains(needle, vm)
1044 ),
1045 ..PySequenceMethods::NOT_IMPLEMENTED
1046 });
1047 &AS_SEQUENCE
1048 }
1049}
1050
1051impl Comparable for PySet {
1052 fn cmp(
1053 zelf: &crate::Py<Self>,
1054 other: &PyObject,
1055 op: PyComparisonOp,
1056 vm: &VirtualMachine,
1057 ) -> PyResult<PyComparisonValue> {
1058 extract_set(other).map_or(Ok(PyComparisonValue::NotImplemented), |other| {
1059 Ok(zelf.inner.compare(other, op, vm)?.into())
1060 })
1061 }
1062}
1063
1064impl Iterable for PySet {
1065 fn iter(zelf: PyRef<Self>, vm: &VirtualMachine) -> PyResult {
1066 Ok(PySetIterator::new(AnySet {
1067 object: zelf.into(),
1068 })
1069 .into_pyobject(vm))
1070 }
1071}
1072
1073impl AsNumber for PySet {
1074 fn as_number() -> &'static PyNumberMethods {
1075 static AS_NUMBER: PyNumberMethods = PyNumberMethods {
1076 subtract: Some(|a, b, vm| {
1079 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1080 return Ok(vm.ctx.not_implemented());
1081 }
1082 if let Some(a) = a.downcast_ref::<PySet>() {
1083 a.__sub__(b.to_owned(), vm).to_pyresult(vm)
1084 } else if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1085 a.__sub__(b.to_owned(), vm)
1087 .map(|r| r.map(|s| PySet { inner: s.inner }))
1088 .to_pyresult(vm)
1089 } else {
1090 Ok(vm.ctx.not_implemented())
1091 }
1092 }),
1093 and: Some(|a, b, vm| {
1094 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1095 return Ok(vm.ctx.not_implemented());
1096 }
1097 if let Some(a) = a.downcast_ref::<PySet>() {
1098 a.__and__(b.to_owned(), vm).to_pyresult(vm)
1099 } else if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1100 a.__and__(b.to_owned(), vm)
1101 .map(|r| r.map(|s| PySet { inner: s.inner }))
1102 .to_pyresult(vm)
1103 } else {
1104 Ok(vm.ctx.not_implemented())
1105 }
1106 }),
1107 xor: Some(|a, b, vm| {
1108 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1109 return Ok(vm.ctx.not_implemented());
1110 }
1111 if let Some(a) = a.downcast_ref::<PySet>() {
1112 a.__xor__(b.to_owned(), vm).to_pyresult(vm)
1113 } else if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1114 a.__xor__(b.to_owned(), vm)
1115 .map(|r| r.map(|s| PySet { inner: s.inner }))
1116 .to_pyresult(vm)
1117 } else {
1118 Ok(vm.ctx.not_implemented())
1119 }
1120 }),
1121 or: Some(|a, b, vm| {
1122 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1123 return Ok(vm.ctx.not_implemented());
1124 }
1125 if let Some(a) = a.downcast_ref::<PySet>() {
1126 a.__or__(b.to_owned(), vm).to_pyresult(vm)
1127 } else if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1128 a.__or__(b.to_owned(), vm)
1129 .map(|r| r.map(|s| PySet { inner: s.inner }))
1130 .to_pyresult(vm)
1131 } else {
1132 Ok(vm.ctx.not_implemented())
1133 }
1134 }),
1135 inplace_subtract: Some(|a, b, vm| {
1136 if let Some(a) = a.downcast_ref::<PySet>() {
1137 PySet::__isub__(a.to_owned(), AnySet::try_from_object(vm, b.to_owned())?, vm)
1138 .to_pyresult(vm)
1139 } else {
1140 Ok(vm.ctx.not_implemented())
1141 }
1142 }),
1143 inplace_and: Some(|a, b, vm| {
1144 if let Some(a) = a.downcast_ref::<PySet>() {
1145 PySet::__iand__(a.to_owned(), AnySet::try_from_object(vm, b.to_owned())?, vm)
1146 .to_pyresult(vm)
1147 } else {
1148 Ok(vm.ctx.not_implemented())
1149 }
1150 }),
1151 inplace_xor: Some(|a, b, vm| {
1152 if let Some(a) = a.downcast_ref::<PySet>() {
1153 PySet::__ixor__(a.to_owned(), AnySet::try_from_object(vm, b.to_owned())?, vm)
1154 .to_pyresult(vm)
1155 } else {
1156 Ok(vm.ctx.not_implemented())
1157 }
1158 }),
1159 inplace_or: Some(|a, b, vm| {
1160 if let Some(a) = a.downcast_ref::<PySet>() {
1161 PySet::__ior__(a.to_owned(), AnySet::try_from_object(vm, b.to_owned())?, vm)
1162 .to_pyresult(vm)
1163 } else {
1164 Ok(vm.ctx.not_implemented())
1165 }
1166 }),
1167 ..PyNumberMethods::NOT_IMPLEMENTED
1168 };
1169 &AS_NUMBER
1170 }
1171}
1172
1173impl Representable for PySet {
1174 #[inline]
1175 fn repr_wtf8(zelf: &crate::Py<Self>, vm: &VirtualMachine) -> PyResult<Wtf8Buf> {
1176 let class = zelf.class();
1177 let borrowed_name = class.name();
1178 let class_name = &*borrowed_name;
1179
1180 if zelf.inner.len() == 0 {
1181 return Ok(Wtf8Buf::from(format!("{class_name}()")));
1182 }
1183
1184 if let Some(_guard) = ReprGuard::enter(vm, zelf.as_object()) {
1185 let name = (class_name != "set").then_some(class_name);
1186 zelf.inner.repr(name, vm)
1187 } else {
1188 Ok(Wtf8Buf::from(format!("{class_name}(...)")))
1189 }
1190 }
1191}
1192
1193impl Constructor for PyFrozenSet {
1194 type Args = crate::function::PositionalIterable;
1195
1196 fn slot_new(cls: PyTypeRef, args: FuncArgs, vm: &VirtualMachine) -> PyResult {
1197 let is_exact_frozenset = cls.is(vm.ctx.types.frozenset_type);
1198 let is_frozenset_init = {
1199 let cls_init = cls
1200 .slots
1201 .init
1202 .load()
1203 .map(|init| crate::types::fn_addr(init));
1204 let frozenset_init = vm
1205 .ctx
1206 .types
1207 .frozenset_type
1208 .slots
1209 .init
1210 .load()
1211 .map(|init| crate::types::fn_addr(init));
1212 cls_init == frozenset_init
1213 };
1214
1215 let iterable_opt = if is_exact_frozenset || is_frozenset_init {
1217 let iterable: crate::function::PositionalIterable = args.bind_for(vm, Self::NAME)?;
1218 let iterable = iterable.iterable;
1219
1220 if is_exact_frozenset
1222 && let OptionalArg::Present(input) = &iterable
1223 && input.class().is(vm.ctx.types.frozenset_type)
1224 {
1225 return Ok(input.clone());
1226 }
1227
1228 iterable
1229 } else {
1230 match &args.args[..] {
1231 [] => OptionalArg::Missing,
1232 [iterable] => OptionalArg::Present(iterable.clone()),
1233 slice => {
1234 return Err(vm.new_arity_type_error(Self::NAME, 0..=1, slice.len()));
1235 }
1236 }
1237 };
1238
1239 let payload = Self::py_new(
1240 &cls,
1241 Self::Args {
1242 iterable: iterable_opt,
1243 },
1244 vm,
1245 )?;
1246
1247 if is_exact_frozenset && payload.inner.len() == 0 {
1249 return Ok(vm.ctx.empty_frozenset.clone().into());
1250 }
1251
1252 payload.into_ref_with_type(vm, cls).map(Into::into)
1253 }
1254
1255 fn py_new(_cls: &Py<PyType>, args: Self::Args, vm: &VirtualMachine) -> PyResult<Self> {
1256 let inner = match args.iterable {
1257 OptionalArg::Present(iterable) => PySetInner::from_object(iterable, vm)?,
1258 OptionalArg::Missing => PySetInner::default(),
1259 };
1260 Ok(Self {
1261 inner,
1262 ..Default::default()
1263 })
1264 }
1265}
1266
1267impl PyFrozenSet {
1268 fn __len__(&self) -> usize {
1269 self.inner.len()
1270 }
1271
1272 pub fn contains(&self, needle: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
1273 self.inner.contains(needle, vm)
1274 }
1275
1276 fn __or__(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyArithmeticValue<Self>> {
1277 if let Ok(set) = AnySet::try_from_object(vm, other) {
1278 Ok(PyArithmeticValue::Implemented(self.op(
1279 set,
1280 PySetInner::union,
1281 vm,
1282 )?))
1283 } else {
1284 Ok(PyArithmeticValue::NotImplemented)
1285 }
1286 }
1287
1288 fn __and__(
1289 &self,
1290 other: PyObjectRef,
1291 vm: &VirtualMachine,
1292 ) -> PyResult<PyArithmeticValue<Self>> {
1293 if let Ok(other) = AnySet::try_from_object(vm, other) {
1294 Ok(PyArithmeticValue::Implemented(Self {
1295 inner: self.inner.intersection(other.into_iterable(vm)?, vm)?,
1296 ..Default::default()
1297 }))
1298 } else {
1299 Ok(PyArithmeticValue::NotImplemented)
1300 }
1301 }
1302
1303 fn __sub__(
1304 &self,
1305 other: PyObjectRef,
1306 vm: &VirtualMachine,
1307 ) -> PyResult<PyArithmeticValue<Self>> {
1308 if let Ok(other) = AnySet::try_from_object(vm, other) {
1309 Ok(PyArithmeticValue::Implemented(Self {
1310 inner: self.inner.difference_new(other.into_iterable(vm)?, vm)?,
1311 ..Default::default()
1312 }))
1313 } else {
1314 Ok(PyArithmeticValue::NotImplemented)
1315 }
1316 }
1317
1318 fn __rsub__(
1319 zelf: PyRef<Self>,
1320 other: PyObjectRef,
1321 vm: &VirtualMachine,
1322 ) -> PyResult<PyArithmeticValue<Self>> {
1323 if let Ok(other) = AnySet::try_from_object(vm, other) {
1324 Ok(PyArithmeticValue::Implemented(Self {
1325 inner: other
1326 .as_inner()
1327 .difference_new(ArgIterable::try_from_object(vm, zelf.into())?, vm)?,
1328 ..Default::default()
1329 }))
1330 } else {
1331 Ok(PyArithmeticValue::NotImplemented)
1332 }
1333 }
1334
1335 fn __xor__(
1336 &self,
1337 other: PyObjectRef,
1338 vm: &VirtualMachine,
1339 ) -> PyResult<PyArithmeticValue<Self>> {
1340 if let Ok(other) = AnySet::try_from_object(vm, other) {
1341 Ok(PyArithmeticValue::Implemented(self.op(
1342 other,
1343 PySetInner::symmetric_difference,
1344 vm,
1345 )?))
1346 } else {
1347 Ok(PyArithmeticValue::NotImplemented)
1348 }
1349 }
1350}
1351
1352#[pyclass(
1353 flags(BASETYPE, _MATCH_SELF, HAS_WEAKREF),
1354 with(
1355 Constructor,
1356 AsSequence,
1357 Hashable,
1358 Comparable,
1359 Iterable,
1360 AsNumber,
1361 Representable
1362 )
1363)]
1364impl Py<PyFrozenSet> {
1365 #[pymethod(coexist)]
1366 fn __contains__(&self, object: PyObjectRef, vm: &VirtualMachine) -> PyResult<bool> {
1367 self.contains(&object, vm)
1368 }
1369
1370 #[pymethod]
1371 fn __sizeof__(&self) -> usize {
1372 core::mem::size_of::<PyFrozenSet>() + self.inner.sizeof()
1373 }
1374
1375 #[pymethod]
1376 fn copy(zelf: PyRef<PyFrozenSet>, vm: &VirtualMachine) -> PyRef<PyFrozenSet> {
1377 if zelf.class().is(vm.ctx.types.frozenset_type) {
1378 zelf
1379 } else {
1380 PyFrozenSet {
1381 inner: zelf.inner.copy(),
1382 ..Default::default()
1383 }
1384 .into_ref(&vm.ctx)
1385 }
1386 }
1387
1388 #[pymethod]
1389 fn union(
1390 &self,
1391 others: PosArgs<ArgIterable, NameOthers>,
1392 vm: &VirtualMachine,
1393 ) -> PyResult<PyFrozenSet> {
1394 self.fold_op(others.into_iter(), PySetInner::union, vm)
1395 }
1396
1397 #[pymethod]
1398 fn intersection(
1399 &self,
1400 others: PosArgs<ArgIterable, NameOthers>,
1401 vm: &VirtualMachine,
1402 ) -> PyResult<PyFrozenSet> {
1403 Ok(PyFrozenSet {
1404 inner: self.inner.intersection_multi(others.into_iter(), vm)?,
1405 ..Default::default()
1406 })
1407 }
1408
1409 #[pymethod]
1410 fn difference(
1411 &self,
1412 others: PosArgs<ArgIterable, NameOthers>,
1413 vm: &VirtualMachine,
1414 ) -> PyResult<PyFrozenSet> {
1415 Ok(PyFrozenSet {
1416 inner: self.inner.difference_multi(others.into_iter(), vm)?,
1417 ..Default::default()
1418 })
1419 }
1420
1421 #[pymethod]
1422 fn symmetric_difference(
1423 &self,
1424 other: ArgIterable,
1425 vm: &VirtualMachine,
1426 ) -> PyResult<PyFrozenSet> {
1427 self.fold_op(
1428 core::iter::once(other),
1429 PySetInner::symmetric_difference,
1430 vm,
1431 )
1432 }
1433
1434 #[pymethod]
1435 fn issubset(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
1436 self.inner.issubset(other, vm)
1437 }
1438
1439 #[pymethod]
1440 fn issuperset(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
1441 self.inner.issuperset(other, vm)
1442 }
1443
1444 #[pymethod]
1445 fn isdisjoint(&self, other: ArgIterable, vm: &VirtualMachine) -> PyResult<bool> {
1446 self.inner.isdisjoint(other, vm)
1447 }
1448
1449 #[pymethod]
1450 fn __reduce__(
1451 zelf: PyRef<PyFrozenSet>,
1452 vm: &VirtualMachine,
1453 ) -> (PyTypeRef, PyTupleRef, Option<PyDictRef>) {
1454 reduce_set(zelf.as_ref(), vm)
1455 }
1456
1457 #[pyclassmethod]
1458 fn __class_getitem__(
1459 cls: PyTypeRef,
1460 object: PyObjectRef,
1461 vm: &VirtualMachine,
1462 ) -> PyResult<PyGenericAlias> {
1463 PyGenericAlias::from_args(cls, object, vm)
1464 }
1465}
1466
1467impl AsSequence for PyFrozenSet {
1468 fn as_sequence() -> &'static PySequenceMethods {
1469 static AS_SEQUENCE: LazyLock<PySequenceMethods> = LazyLock::new(|| PySequenceMethods {
1470 length: atomic_func!(|seq, _vm| Ok(PyFrozenSet::sequence_downcast(seq).__len__())),
1471 contains: atomic_func!(
1472 |seq, needle, vm| PyFrozenSet::sequence_downcast(seq).contains(needle, vm)
1473 ),
1474 ..PySequenceMethods::NOT_IMPLEMENTED
1475 });
1476 &AS_SEQUENCE
1477 }
1478}
1479
1480impl Hashable for PyFrozenSet {
1481 #[inline]
1482 fn hash(zelf: &crate::Py<Self>, _vm: &VirtualMachine) -> PyResult<PyHash> {
1483 let hash = match zelf.hash.load(Ordering::Relaxed) {
1484 hash::SENTINEL => {
1485 let hash = zelf.inner.hash();
1486 match Radium::compare_exchange(
1487 &zelf.hash,
1488 hash::SENTINEL,
1489 hash::fix_sentinel(hash),
1490 Ordering::Relaxed,
1491 Ordering::Relaxed,
1492 ) {
1493 Ok(_) => hash,
1494 Err(prev_stored) => prev_stored,
1495 }
1496 }
1497 hash => hash,
1498 };
1499 Ok(hash)
1500 }
1501}
1502
1503impl Comparable for PyFrozenSet {
1504 fn cmp(
1505 zelf: &crate::Py<Self>,
1506 other: &PyObject,
1507 op: PyComparisonOp,
1508 vm: &VirtualMachine,
1509 ) -> PyResult<PyComparisonValue> {
1510 extract_set(other).map_or(Ok(PyComparisonValue::NotImplemented), |other| {
1511 Ok(zelf.inner.compare(other, op, vm)?.into())
1512 })
1513 }
1514}
1515
1516impl Iterable for PyFrozenSet {
1517 fn iter(zelf: PyRef<Self>, vm: &VirtualMachine) -> PyResult {
1518 Ok(PySetIterator::new(AnySet {
1519 object: zelf.into(),
1520 })
1521 .into_pyobject(vm))
1522 }
1523}
1524
1525impl AsNumber for PyFrozenSet {
1526 fn as_number() -> &'static PyNumberMethods {
1527 static AS_NUMBER: PyNumberMethods = PyNumberMethods {
1528 subtract: Some(|a, b, vm| {
1531 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1532 return Ok(vm.ctx.not_implemented());
1533 }
1534 if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1535 a.__sub__(b.to_owned(), vm).to_pyresult(vm)
1536 } else if let Some(a) = a.downcast_ref::<PySet>() {
1537 a.__sub__(b.to_owned(), vm).to_pyresult(vm)
1539 } else {
1540 Ok(vm.ctx.not_implemented())
1541 }
1542 }),
1543 and: Some(|a, b, vm| {
1544 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1545 return Ok(vm.ctx.not_implemented());
1546 }
1547 if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1548 a.__and__(b.to_owned(), vm).to_pyresult(vm)
1549 } else if let Some(a) = a.downcast_ref::<PySet>() {
1550 a.__and__(b.to_owned(), vm).to_pyresult(vm)
1551 } else {
1552 Ok(vm.ctx.not_implemented())
1553 }
1554 }),
1555 xor: Some(|a, b, vm| {
1556 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1557 return Ok(vm.ctx.not_implemented());
1558 }
1559 if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1560 a.__xor__(b.to_owned(), vm).to_pyresult(vm)
1561 } else if let Some(a) = a.downcast_ref::<PySet>() {
1562 a.__xor__(b.to_owned(), vm).to_pyresult(vm)
1563 } else {
1564 Ok(vm.ctx.not_implemented())
1565 }
1566 }),
1567 or: Some(|a, b, vm| {
1568 if !AnySet::check(a, vm) || !AnySet::check(b, vm) {
1569 return Ok(vm.ctx.not_implemented());
1570 }
1571 if let Some(a) = a.downcast_ref::<PyFrozenSet>() {
1572 a.__or__(b.to_owned(), vm).to_pyresult(vm)
1573 } else if let Some(a) = a.downcast_ref::<PySet>() {
1574 a.__or__(b.to_owned(), vm).to_pyresult(vm)
1575 } else {
1576 Ok(vm.ctx.not_implemented())
1577 }
1578 }),
1579 ..PyNumberMethods::NOT_IMPLEMENTED
1580 };
1581 &AS_NUMBER
1582 }
1583}
1584
1585impl Representable for PyFrozenSet {
1586 #[inline]
1587 fn repr_wtf8(zelf: &crate::Py<Self>, vm: &VirtualMachine) -> PyResult<Wtf8Buf> {
1588 let inner = &zelf.inner;
1589 let class = zelf.class();
1590 let class_name = class.name();
1591 if inner.len() == 0 {
1592 return Ok(Wtf8Buf::from(format!("{class_name}()")));
1593 }
1594 if let Some(_guard) = ReprGuard::enter(vm, zelf.as_object()) {
1595 inner.repr(Some(&class_name), vm)
1596 } else {
1597 Ok(Wtf8Buf::from(format!("{class_name}(...)")))
1598 }
1599 }
1600}
1601
1602struct AnySet {
1603 object: PyObjectRef,
1604}
1605
1606impl Borrow<PyObject> for AnySet {
1607 #[inline(always)]
1608 fn borrow(&self) -> &PyObject {
1609 &self.object
1610 }
1611}
1612
1613impl AnySet {
1614 fn check(obj: &PyObject, vm: &VirtualMachine) -> bool {
1617 let ctx = &vm.ctx;
1618 obj.fast_isinstance(ctx.types.set_type) || obj.fast_isinstance(ctx.types.frozenset_type)
1619 }
1620
1621 fn into_iterable(self, vm: &VirtualMachine) -> PyResult<ArgIterable> {
1622 self.object.try_into_value(vm)
1623 }
1624
1625 fn into_iterable_iter(
1626 self,
1627 vm: &VirtualMachine,
1628 ) -> PyResult<impl core::iter::Iterator<Item = ArgIterable>> {
1629 Ok(core::iter::once(self.into_iterable(vm)?))
1630 }
1631
1632 fn as_inner(&self) -> &PySetInner {
1633 match_class!(match self.object.as_object() {
1634 ref set @ PySet => &set.inner,
1635 ref frozen @ PyFrozenSet => &frozen.inner,
1636 _ => unreachable!("AnySet is always PySet or PyFrozenSet"), })
1638 }
1639}
1640
1641impl TryFromObject for AnySet {
1642 fn try_from_object(vm: &VirtualMachine, obj: PyObjectRef) -> PyResult<Self> {
1643 let class = obj.class();
1644 if class.fast_issubclass(vm.ctx.types.set_type)
1645 || class.fast_issubclass(vm.ctx.types.frozenset_type)
1646 {
1647 Ok(Self { object: obj })
1648 } else {
1649 Err(vm.new_type_error(format!("{class} is not a subtype of set or frozenset")))
1650 }
1651 }
1652}
1653
1654#[pyclass(module = false, name = "set_iterator")]
1655pub(crate) struct PySetIterator {
1656 size: DictSize,
1657 changed: PyAtomic<bool>,
1661 internal: PyMutex<PositionIterInternal<AnySet>>,
1662}
1663
1664impl fmt::Debug for PySetIterator {
1665 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1666 f.write_str("set_iterator")
1668 }
1669}
1670
1671impl PyPayload for PySetIterator {
1672 #[inline]
1673 fn class(ctx: &Context) -> &'static Py<PyType> {
1674 ctx.types.set_iterator_type
1675 }
1676}
1677
1678impl PySetIterator {
1679 fn new(set: AnySet) -> Self {
1680 Self {
1681 size: set.as_inner().content.size(),
1682 changed: Radium::new(false),
1683 internal: PyMutex::new(PositionIterInternal::new(set, 0)),
1684 }
1685 }
1686}
1687
1688#[pyclass(flags(DISALLOW_INSTANTIATION), with(IterNext, Iterable))]
1689impl Py<PySetIterator> {
1690 #[pymethod]
1691 fn __length_hint__(&self) -> usize {
1692 if self.changed.load(Ordering::Relaxed) {
1696 return 0;
1697 }
1698 self.internal.lock().length_hint(|set| {
1699 if set.as_inner().content.size() == self.size {
1700 self.size.entries_size
1701 } else {
1702 0
1703 }
1704 })
1705 }
1706
1707 #[pymethod]
1708 fn __reduce__(
1709 zelf: PyRef<PySetIterator>,
1710 vm: &VirtualMachine,
1711 ) -> PyResult<(PyObjectRef, (PyObjectRef,))> {
1712 let internal = zelf.internal.lock();
1713 Ok((
1714 builtins_iter(vm)?,
1715 (vm.ctx
1716 .new_list(match &internal.status {
1717 IterStatus::Exhausted => vec![],
1718 IterStatus::Active(set) => set
1719 .as_inner()
1720 .content
1721 .keys()
1722 .into_iter()
1723 .skip(internal.position)
1724 .collect(),
1725 })
1726 .into(),),
1727 ))
1728 }
1729}
1730
1731impl SelfIter for PySetIterator {}
1732impl IterNext for PySetIterator {
1733 fn next(zelf: &crate::Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
1734 locked_step(&zelf.internal, |internal| {
1735 let IterStatus::Active(set) = &internal.status else {
1736 return (Ok(PyIterReturn::StopIteration(None)), None);
1737 };
1738 let mutated = || vm.new_runtime_error("Set changed size during iteration");
1739 if zelf.changed.load(Ordering::Relaxed) {
1740 return (Err(mutated()), None);
1743 }
1744 let entry = set.as_inner().content.next_entry_checked(
1745 internal.position,
1746 &zelf.size,
1747 |key, ()| key.to_owned(),
1748 );
1749 match entry {
1750 Err(crate::dict_inner::DictChanged) => {
1751 zelf.changed.store(true, Ordering::Relaxed);
1752 (Err(mutated()), None)
1753 }
1754 Ok(Some((position, key))) => {
1755 internal.position = position;
1756 (Ok(PyIterReturn::Return(key)), None)
1757 }
1758 Ok(None) => (Ok(PyIterReturn::StopIteration(None)), internal.exhaust()),
1759 }
1760 })
1761 }
1762}
1763
1764fn vectorcall_set(
1765 zelf_obj: &PyObject,
1766 args: Vec<PyObjectRef>,
1767 nargs: usize,
1768 kwnames: Option<&[PyObjectRef]>,
1769 vm: &VirtualMachine,
1770) -> PyResult {
1771 let zelf: &Py<PyType> = zelf_obj.downcast_ref().unwrap();
1772 let obj = PySet::default().into_ref_with_type(vm, zelf.to_owned())?;
1773 let func_args = FuncArgs::from_vectorcall_owned(args, nargs, kwnames);
1774 PySet::slot_init(obj.as_object(), func_args, vm)?;
1775 Ok(obj.into())
1776}
1777
1778fn vectorcall_frozenset(
1779 zelf_obj: &PyObject,
1780 args: Vec<PyObjectRef>,
1781 nargs: usize,
1782 kwnames: Option<&[PyObjectRef]>,
1783 vm: &VirtualMachine,
1784) -> PyResult {
1785 let zelf: &Py<PyType> = zelf_obj.downcast_ref().unwrap();
1786 let func_args = FuncArgs::from_vectorcall_owned(args, nargs, kwnames);
1787 (zelf.slots.new.load().unwrap())(zelf.to_owned(), func_args, vm)
1788}
1789
1790pub(crate) fn init(context: &'static Context) {
1791 PySet::extend_class(context, context.types.set_type);
1792 context
1793 .types
1794 .set_type
1795 .slots
1796 .vectorcall
1797 .store(Some(vectorcall_set));
1798
1799 PyFrozenSet::extend_class(context, context.types.frozenset_type);
1800 context
1801 .types
1802 .frozenset_type
1803 .slots
1804 .vectorcall
1805 .store(Some(vectorcall_frozenset));
1806
1807 PySetIterator::extend_class(context, context.types.set_iterator_type);
1808}