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 hashable_args: Option<PyRef<PyFrozenSet>>,
25 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 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 #[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 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 if args_tuple.as_slice().is_empty() {
194 return Err(vm.new_type_error("Cannot create empty Union"));
195 }
196
197 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 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(¶meters, 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 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 let annotationlib = vm.import("annotationlib", 0)?;
280 let forwardref_cls = annotationlib.get_attr("ForwardRef", vm)?;
281 forwardref_cls.call((arg,), vm)
282}
283
284struct UnionComponents {
286 args: PyTupleRef,
288 hashable_args: Option<PyRef<PyFrozenSet>>,
290 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 let mut new_args: Vec<PyObjectRef> = Vec::with_capacity(args.as_slice().len());
305
306 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 match arg.hash(vm) {
314 Ok(_) => {
315 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 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 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 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 if zelf.args.as_slice().len() != other.args.as_slice().len() {
435 return Ok(PyComparisonValue::Implemented(false));
436 }
437
438 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 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), }
463 }
464 if !found {
465 return Ok(PyComparisonValue::Implemented(false));
466 }
467 }
468
469 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), }
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 let Some(ref unhashable_args) = zelf.unhashable_args {
497 let n = unhashable_args.as_slice().len();
498 for arg in unhashable_args.as_slice() {
500 arg.hash(vm)?;
501 }
502 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 let Some(ref hashable_args) = zelf.hashable_args {
513 return PyFrozenSet::hash(hashable_args, vm);
514 }
515
516 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}