Skip to main content

rustpython_vm/builtins/
union.rs

1use super::{genericalias, type_};
2use crate::common::lock::LazyLock;
3use crate::{
4    AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
5    atomic_func,
6    builtins::{PyFrozenSet, PySet, PyStr, PyTuple, PyTupleRef, PyType},
7    class::PyClassImpl,
8    common::hash,
9    convert::ToPyObject,
10    function::PyComparisonValue,
11    protocol::{PyMappingMethods, PyNumberMethods},
12    stdlib::_typing::{TypeAliasType, call_typing_func_object},
13    types::{AsMapping, AsNumber, Comparable, GetAttr, Hashable, PyComparisonOp, Representable},
14};
15use alloc::fmt;
16
17const CLS_ATTRS: &[&str] = &["__module__"];
18
19#[pyclass(module = "typing", name = "Union", traverse)]
20pub struct PyUnion {
21    #[pymember(name = "__args__")]
22    args: PyTupleRef,
23    /// Frozenset of hashable args, or None if all args were hashable
24    hashable_args: Option<PyRef<PyFrozenSet>>,
25    /// Tuple of initially unhashable args, or None if all args were hashable
26    unhashable_args: Option<PyTupleRef>,
27    parameters: PyTupleRef,
28}
29
30impl fmt::Debug for PyUnion {
31    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
32        f.write_str("UnionObject")
33    }
34}
35
36impl PyPayload for PyUnion {
37    #[inline]
38    fn class(ctx: &Context) -> &'static Py<PyType> {
39        ctx.types.union_type
40    }
41}
42
43impl PyUnion {
44    /// Create a new union from dedup result (internal use)
45    fn from_components(result: UnionComponents, vm: &VirtualMachine) -> PyResult<Self> {
46        let parameters = make_parameters(&result.args, vm)?;
47        Ok(Self {
48            args: result.args,
49            hashable_args: result.hashable_args,
50            unhashable_args: result.unhashable_args,
51            parameters,
52        })
53    }
54
55    /// Direct access to args field (_Py_union_args)
56    #[inline]
57    #[must_use]
58    pub fn args(&self) -> &Py<PyTuple> {
59        &self.args
60    }
61
62    fn repr(&self, vm: &VirtualMachine) -> PyResult<String> {
63        fn repr_item(obj: &PyObject, vm: &VirtualMachine) -> PyResult<String> {
64            if obj.is(vm.ctx.types.none_type) {
65                return Ok("None".to_string());
66            }
67
68            if vm
69                .get_attribute_opt(obj, identifier!(vm, __origin__))?
70                .is_some()
71                && vm
72                    .get_attribute_opt(obj, identifier!(vm, __args__))?
73                    .is_some()
74            {
75                return Ok(obj.repr(vm)?.to_string());
76            }
77
78            match (
79                vm.get_attribute_opt(obj, identifier!(vm, __qualname__))?
80                    .and_then(|o| o.downcast_ref::<PyStr>().map(|n| n.to_string())),
81                vm.get_attribute_opt(obj, identifier!(vm, __module__))?
82                    .and_then(|o| o.downcast_ref::<PyStr>().map(|m| m.to_string())),
83            ) {
84                (None, _) | (_, None) => Ok(obj.repr(vm)?.to_string()),
85                (Some(qualname), Some(module)) => Ok(if module == "builtins" {
86                    qualname
87                } else {
88                    format!("{module}.{qualname}")
89                }),
90            }
91        }
92
93        Ok(self
94            .args
95            .as_slice()
96            .iter()
97            .map(|o| repr_item(o, vm))
98            .collect::<PyResult<Vec<_>>>()?
99            .join(" | "))
100    }
101}
102
103impl PyUnion {
104    fn __or__(zelf: PyObjectRef, other: PyObjectRef, vm: &VirtualMachine) -> PyResult {
105        type_::or_(zelf, other, vm)
106    }
107}
108
109#[pyclass(
110    flags(DISALLOW_INSTANTIATION, HAS_WEAKREF),
111    with(Hashable, Comparable, AsMapping, AsNumber, Representable)
112)]
113impl Py<PyUnion> {
114    #[pygetset]
115    fn __name__(&self, vm: &VirtualMachine) -> PyObjectRef {
116        vm.ctx.new_str("Union").into()
117    }
118
119    #[pygetset]
120    fn __qualname__(&self, vm: &VirtualMachine) -> PyObjectRef {
121        vm.ctx.new_str("Union").into()
122    }
123
124    #[pygetset]
125    fn __origin__(&self, vm: &VirtualMachine) -> PyObjectRef {
126        vm.ctx.types.union_type.to_owned().into()
127    }
128
129    #[pygetset]
130    fn __parameters__(&self) -> PyObjectRef {
131        self.parameters.clone().into()
132    }
133
134    #[pymethod]
135    fn __instancecheck__(
136        zelf: PyRef<PyUnion>,
137        obj: PyObjectRef,
138        vm: &VirtualMachine,
139    ) -> PyResult<bool> {
140        if zelf
141            .args
142            .as_slice()
143            .iter()
144            .any(|x| x.class().is(vm.ctx.types.generic_alias_type))
145        {
146            Err(vm.new_type_error("isinstance() argument 2 cannot be a parameterized generic"))
147        } else {
148            obj.is_instance(zelf.args.as_object(), vm)
149        }
150    }
151
152    #[pymethod]
153    fn __subclasscheck__(
154        zelf: PyRef<PyUnion>,
155        obj: PyObjectRef,
156        vm: &VirtualMachine,
157    ) -> PyResult<bool> {
158        if zelf
159            .args
160            .as_slice()
161            .iter()
162            .any(|x| x.class().is(vm.ctx.types.generic_alias_type))
163        {
164            Err(vm.new_type_error("issubclass() argument 2 cannot be a parameterized generic"))
165        } else {
166            obj.is_subclass(zelf.args.as_object(), vm)
167        }
168    }
169
170    #[pymethod]
171    fn __mro_entries__(
172        zelf: PyRef<PyUnion>,
173        _object: PyObjectRef,
174        vm: &VirtualMachine,
175    ) -> PyResult {
176        Err(vm.new_type_error(format!("Cannot subclass {}", zelf.repr(vm)?)))
177    }
178
179    #[pyclassmethod]
180    fn __class_getitem__(
181        _cls: crate::builtins::PyTypeRef,
182        object: PyObjectRef,
183        vm: &VirtualMachine,
184    ) -> PyResult {
185        // Convert args to tuple if not already
186        let args_tuple = if let Some(tuple) = object.downcast_ref::<PyTuple>() {
187            tuple.to_owned()
188        } else {
189            PyTuple::new_ref(vec![object], &vm.ctx)
190        };
191
192        // Check for empty union
193        if args_tuple.as_slice().is_empty() {
194            return Err(vm.new_type_error("Cannot create empty Union"));
195        }
196
197        // Create union using make_union to properly handle None -> NoneType conversion
198        make_union(&args_tuple, vm)
199    }
200}
201
202fn is_unionable(obj: &PyObject, vm: &VirtualMachine) -> bool {
203    let cls = obj.class();
204    cls.is(vm.ctx.types.none_type)
205        || obj.downcastable::<PyType>()
206        || cls.fast_issubclass(vm.ctx.types.generic_alias_type)
207        || cls.is(vm.ctx.types.union_type)
208        || obj.downcast_ref::<TypeAliasType>().is_some()
209}
210
211fn type_check(arg: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
212    // Fast path to avoid calling into typing.py
213    if is_unionable(&arg, vm) {
214        return Ok(arg);
215    }
216    let message_str: PyObjectRef = vm
217        .ctx
218        .new_str("Union[arg, ...]: each arg must be a type.")
219        .into();
220    call_typing_func_object(vm, "_type_check", (arg, message_str))
221}
222
223fn has_union_operands(a: &PyObject, b: &PyObject, vm: &VirtualMachine) -> bool {
224    let union_type = vm.ctx.types.union_type;
225    a.class().is(union_type) || b.class().is(union_type)
226}
227
228pub(crate) fn or_op(zelf: PyObjectRef, other: PyObjectRef, vm: &VirtualMachine) -> PyResult {
229    if !has_union_operands(&zelf, &other, vm)
230        && (!is_unionable(&zelf, vm) || !is_unionable(&other, vm))
231    {
232        return Ok(vm.ctx.not_implemented());
233    }
234
235    let left = type_check(zelf, vm)?;
236    let right = type_check(other, vm)?;
237    let tuple = PyTuple::new_ref(vec![left, right], &vm.ctx);
238    make_union(&tuple, vm)
239}
240
241fn make_parameters(args: &Py<PyTuple>, vm: &VirtualMachine) -> PyResult<PyTupleRef> {
242    let parameters = genericalias::make_parameters(args, vm)?;
243    let result = dedup_and_flatten_args(&parameters, vm)?;
244    Ok(result.args)
245}
246
247fn flatten_args(args: &Py<PyTuple>, vm: &VirtualMachine) -> PyTupleRef {
248    let mut total_args = 0;
249    for arg in args {
250        if let Some(pyref) = arg.downcast_ref::<PyUnion>() {
251            total_args += pyref.args.as_slice().len();
252        } else {
253            total_args += 1;
254        };
255    }
256
257    let mut flattened_args = Vec::with_capacity(total_args);
258    for arg in args {
259        if let Some(pyref) = arg.downcast_ref::<PyUnion>() {
260            flattened_args.extend(pyref.args.as_slice().iter().cloned());
261        } else if vm.is_none(arg) {
262            flattened_args.push(vm.ctx.types.none_type.to_owned().into());
263        } else if arg.downcast_ref::<PyStr>().is_some() {
264            // Convert string to ForwardRef
265            match string_to_forwardref(arg.clone(), vm) {
266                Ok(fr) => flattened_args.push(fr),
267                Err(_) => flattened_args.push(arg.clone()),
268            }
269        } else {
270            flattened_args.push(arg.clone());
271        };
272    }
273
274    PyTuple::new_ref(flattened_args, &vm.ctx)
275}
276
277fn string_to_forwardref(arg: PyObjectRef, vm: &VirtualMachine) -> PyResult {
278    // Import annotationlib.ForwardRef and create a ForwardRef
279    let annotationlib = vm.import("annotationlib", 0)?;
280    let forwardref_cls = annotationlib.get_attr("ForwardRef", vm)?;
281    forwardref_cls.call((arg,), vm)
282}
283
284/// Components for creating a PyUnion after deduplication
285struct UnionComponents {
286    /// All unique args in order
287    args: PyTupleRef,
288    /// Frozenset of hashable args (for fast equality comparison)
289    hashable_args: Option<PyRef<PyFrozenSet>>,
290    /// Tuple of unhashable args at creation time (for hash error message)
291    unhashable_args: Option<PyTupleRef>,
292}
293
294fn dedup_and_flatten_args(args: &Py<PyTuple>, vm: &VirtualMachine) -> PyResult<UnionComponents> {
295    let args = flatten_args(args, vm);
296
297    // Use set-based deduplication like CPython:
298    // - For hashable elements: use Python's set semantics (hash + equality)
299    // - For unhashable elements: use equality comparison
300    //
301    // This avoids calling __eq__ when hashes differ, so `int | BadType`
302    // doesn't raise even if BadType.__eq__ raises.
303
304    let mut new_args: Vec<PyObjectRef> = Vec::with_capacity(args.as_slice().len());
305
306    // Track hashable elements using a Python set (uses hash + equality)
307    let hashable_set = PySet::default().into_ref(&vm.ctx);
308    let mut hashable_list: Vec<PyObjectRef> = Vec::new();
309    let mut unhashable_list: Vec<PyObjectRef> = Vec::new();
310
311    for arg in &*args {
312        // Try to hash the element first
313        match arg.hash(vm) {
314            Ok(_) => {
315                // Element is hashable - use set for deduplication
316                // Set membership uses hash first, then equality only if hashes match
317                let contains = vm
318                    .call_method(hashable_set.as_ref(), "__contains__", (arg.clone(),))
319                    .and_then(|r| r.try_to_bool(vm))?;
320                if !contains {
321                    hashable_set.add(arg.clone(), vm)?;
322                    hashable_list.push(arg.clone());
323                    new_args.push(arg.clone());
324                }
325            }
326            Err(_) => {
327                // Element is unhashable - use equality comparison
328                let mut is_duplicate = false;
329                for existing in &unhashable_list {
330                    match existing.rich_compare_bool(arg, PyComparisonOp::Eq, vm) {
331                        Ok(true) => {
332                            is_duplicate = true;
333                            break;
334                        }
335                        Ok(false) => continue,
336                        Err(e) => return Err(e),
337                    }
338                }
339                if !is_duplicate {
340                    unhashable_list.push(arg.clone());
341                    new_args.push(arg.clone());
342                }
343            }
344        }
345    }
346
347    new_args.shrink_to_fit();
348
349    // Create hashable_args frozenset if there are hashable elements
350    let hashable_args = if !hashable_list.is_empty() {
351        Some(PyFrozenSet::from_iter(vm, hashable_list)?.into_ref(&vm.ctx))
352    } else {
353        None
354    };
355
356    // Create unhashable_args tuple if there are unhashable elements
357    let unhashable_args = if !unhashable_list.is_empty() {
358        Some(PyTuple::new_ref(unhashable_list, &vm.ctx))
359    } else {
360        None
361    };
362
363    Ok(UnionComponents {
364        args: PyTuple::new_ref(new_args, &vm.ctx),
365        hashable_args,
366        unhashable_args,
367    })
368}
369
370pub fn make_union(args: &Py<PyTuple>, vm: &VirtualMachine) -> PyResult {
371    let result = dedup_and_flatten_args(args, vm)?;
372    Ok(match result.args.as_slice().len() {
373        1 => result.args.as_slice()[0].to_owned(),
374        _ => PyUnion::from_components(result, vm)?.to_pyobject(vm),
375    })
376}
377
378impl PyUnion {
379    fn getitem(zelf: &Py<Self>, needle: PyObjectRef, vm: &VirtualMachine) -> PyResult {
380        let new_args = genericalias::subs_parameters(
381            zelf.as_object(),
382            &zelf.args,
383            &zelf.parameters,
384            needle,
385            vm,
386        )?;
387
388        Ok(if new_args.as_slice().is_empty() {
389            make_union(&new_args, vm)?
390        } else {
391            let mut tmp = new_args.as_slice()[0].to_owned();
392            for arg in new_args.as_slice().iter().skip(1) {
393                tmp = vm._or(&tmp, arg)?;
394            }
395            tmp
396        })
397    }
398}
399
400impl AsMapping for PyUnion {
401    fn as_mapping() -> &'static PyMappingMethods {
402        static AS_MAPPING: LazyLock<PyMappingMethods> = LazyLock::new(|| PyMappingMethods {
403            subscript: atomic_func!(|mapping, needle, vm| {
404                let zelf = PyUnion::mapping_downcast(mapping);
405                PyUnion::getitem(zelf, needle.to_owned(), vm)
406            }),
407            ..PyMappingMethods::NOT_IMPLEMENTED
408        });
409        &AS_MAPPING
410    }
411}
412
413impl AsNumber for PyUnion {
414    fn as_number() -> &'static PyNumberMethods {
415        static AS_NUMBER: PyNumberMethods = PyNumberMethods {
416            or: Some(|a, b, vm| PyUnion::__or__(a.to_owned(), b.to_owned(), vm)),
417            ..PyNumberMethods::NOT_IMPLEMENTED
418        };
419        &AS_NUMBER
420    }
421}
422
423impl Comparable for PyUnion {
424    fn cmp(
425        zelf: &Py<Self>,
426        other: &PyObject,
427        op: PyComparisonOp,
428        vm: &VirtualMachine,
429    ) -> PyResult<PyComparisonValue> {
430        op.eq_only(|| {
431            let other = class_or_notimplemented!(Self, other);
432
433            // Check if lengths are equal
434            if zelf.args.as_slice().len() != other.args.as_slice().len() {
435                return Ok(PyComparisonValue::Implemented(false));
436            }
437
438            // Fast path: if both unions have all hashable args, compare frozensets directly
439            // Always use Eq here since eq_only handles Ne by negating the result
440            if zelf.unhashable_args.is_none()
441                && other.unhashable_args.is_none()
442                && let (Some(a), Some(b)) = (&zelf.hashable_args, &other.hashable_args)
443            {
444                let eq = a
445                    .as_object()
446                    .rich_compare_bool(b.as_object(), PyComparisonOp::Eq, vm)?;
447                return Ok(PyComparisonValue::Implemented(eq));
448            }
449
450            // Slow path: O(n^2) nested loop comparison for unhashable elements
451            // Check if all elements in zelf.args are in other.args
452            for arg_a in &*zelf.args {
453                let mut found = false;
454                for arg_b in &*other.args {
455                    match arg_a.rich_compare_bool(arg_b, PyComparisonOp::Eq, vm) {
456                        Ok(true) => {
457                            found = true;
458                            break;
459                        }
460                        Ok(false) => continue,
461                        Err(e) => return Err(e), // Propagate comparison errors
462                    }
463                }
464                if !found {
465                    return Ok(PyComparisonValue::Implemented(false));
466                }
467            }
468
469            // Check if all elements in other.args are in zelf.args (for symmetry)
470            for arg_b in &*other.args {
471                let mut found = false;
472                for arg_a in &*zelf.args {
473                    match arg_b.rich_compare_bool(arg_a, PyComparisonOp::Eq, vm) {
474                        Ok(true) => {
475                            found = true;
476                            break;
477                        }
478                        Ok(false) => continue,
479                        Err(e) => return Err(e), // Propagate comparison errors
480                    }
481                }
482                if !found {
483                    return Ok(PyComparisonValue::Implemented(false));
484                }
485            }
486
487            Ok(PyComparisonValue::Implemented(true))
488        })
489    }
490}
491
492impl Hashable for PyUnion {
493    #[inline]
494    fn hash(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<hash::PyHash> {
495        // If there are any unhashable args from creation time, the union is unhashable
496        if let Some(ref unhashable_args) = zelf.unhashable_args {
497            let n = unhashable_args.as_slice().len();
498            // Try to hash each previously unhashable arg to get an error
499            for arg in unhashable_args.as_slice() {
500                arg.hash(vm)?;
501            }
502            // All previously unhashable args somehow became hashable
503            // But still raise an error to maintain consistent hashing
504            return Err(vm.new_type_error(format!(
505                "union contains {} unhashable element{}",
506                n,
507                if n > 1 { "s" } else { "" }
508            )));
509        }
510
511        // If we have a stored frozenset of hashable args, use that
512        if let Some(ref hashable_args) = zelf.hashable_args {
513            return PyFrozenSet::hash(hashable_args, vm);
514        }
515
516        // Fallback: compute hash from args
517        let mut args_to_hash = Vec::new();
518        for arg in &*zelf.args {
519            match arg.hash(vm) {
520                Ok(_) => args_to_hash.push(arg.clone()),
521                Err(e) => return Err(e),
522            }
523        }
524        let set = PyFrozenSet::from_iter(vm, args_to_hash)?;
525        PyFrozenSet::hash(&set.into_ref(&vm.ctx), vm)
526    }
527}
528
529impl GetAttr for PyUnion {
530    fn getattro(zelf: &Py<Self>, attr: &Py<PyStr>, vm: &VirtualMachine) -> PyResult {
531        for &exc in CLS_ATTRS {
532            if *exc == attr.to_string() {
533                return zelf.as_object().generic_getattr(attr, vm);
534            }
535        }
536        zelf.as_object().get_attr(attr, vm)
537    }
538}
539
540impl Representable for PyUnion {
541    #[inline]
542    fn repr_str(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<String> {
543        zelf.repr(vm)
544    }
545}
546
547pub(crate) fn init(context: &'static Context) {
548    let union_type = &context.types.union_type;
549    PyUnion::extend_class(context, union_type);
550}