Skip to main content

rustpython_vm/builtins/
set.rs

1/*
2 * Builtin set type with a sequence of unique items.
3 */
4use 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    // Also used by ssl.rs windows.
102    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        // FIXME: empty set check
111        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        // TODO: implement more detailed, non-recursive Debug formatter
151        f.write_str("set")
152    }
153}
154
155impl fmt::Debug for PyFrozenSet {
156    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
157        // TODO: implement more detailed, non-recursive Debug formatter
158        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        // FIXME(discord9): Rc means shared ref, so should it be traced?
185        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    /// Build a set from an arbitrary object, reusing stored hashes when the
203    /// source is a set/frozenset/dict.
204    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    /// Elements of `obj` with their stored hashes, or `None` if `obj` keeps
211    /// none and must be iterated generically. Mirrors the `PyAnySet_Check` /
212    /// `PyDict_CheckExact` fast paths in CPython's `set_update_internal`.
213    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    /// Look up a key whose hash is already known, without a frozenset retry.
284    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        // Scanning the left side avoids visiting and retaining a much larger
368        // exclusion collection. Match CPython's comparison direction as well.
369        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    // The dict-view caller supplies a private working table.
386    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            // the source is already duplicate-free
414            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        // We want to remove duplicates in other
423        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    /// [`Self::add`] with a known hash.
499    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        // TODO: should be pop_front, but that requires rearranging every index
526        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        // check AnySet
536        if let Ok(any_set) = AnySet::try_from_object(vm, iterable.to_owned()) {
537            self.merge_set(any_set, vm)
538        // check Dict
539        } else if let Ok(dict) = iterable.to_owned().downcast_exact::<PyDict>(vm) {
540            self.merge_dict(&dict, vm)
541        } else {
542            // add iterable that is not AnySet or Dict
543            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                // Build the intersection first: besides bounding the work by
588                // our size, this preserves CPython's equality/error ordering.
589                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                // the source is already duplicate-free
620                for (item, hash) in elements {
621                    self.content
622                        .delete_or_insert_known_hash(vm, &item, hash, ())?;
623                }
624                continue;
625            }
626            // We want to remove duplicates in iterable
627            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    // Run operation, on failure, if item is a set/set subclass, convert it
648    // into a frozenset and try the operation again. Propagates original error
649    // on failure to convert and restores item in KeyError on failure (remove).
650    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                // Keep original error around.
662                .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                    // If operation raised KeyError, report original set (set.remove)
673                    .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
712/// Elements of `obj` with their stored hashes, or `None` unless `obj` is exactly
713/// a `set` or `frozenset` — `PyAnySet_CheckExact`, where [`extract_set`] is the
714/// subclass-inclusive `PyAnySet_Check`.
715pub(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            // Binary ops check both operands are sets (like CPython's set_sub, etc.)
1077            // This is needed because __rsub__ swaps operands: a.__rsub__(b) calls subtract(b, a)
1078            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                    // When called via __rsub__, a might be PyFrozenSet
1086                    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        // Optimizations for exact frozenset type
1216        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            // Return exact frozenset as-is
1221            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        // Return empty frozenset singleton
1248        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            // Binary ops check both operands are sets (like CPython's set_sub, etc.)
1529            // __rsub__ swaps operands. Result type follows first operand's type.
1530            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                    // When called via __rsub__, a might be PySet - return set (not frozenset)
1538                    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    /// Check if object is a set or frozenset (including subclasses)
1615    /// Equivalent to CPython's PyAnySet_Check
1616    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"), // should not be called.
1637        })
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    /// Whether the set was found to have changed, which `setiter_iternext()`
1658    /// records by writing a size no set can have. Sticky: what it makes the
1659    /// iterator answer, it answers from then on.
1660    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        // TODO: implement more detailed, non-recursive Debug formatter
1667        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        // `setiter_len()` answers for a set it can no longer walk with nothing,
1693        // comparing the size it captured against the set's own every time it is
1694        // asked.
1695        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                // The set is not looked at again once it has been found to
1741                // change: an iterator that has raised keeps raising.
1742                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}