Skip to main content

rustpython_vm/
exception_group.rs

1//! ExceptionGroup implementation for Python 3.11+
2//!
3//! This module implements BaseExceptionGroup and ExceptionGroup with multiple inheritance support.
4
5use crate::builtins::{PyList, PyStrRef, PyTuple, PyTupleRef, PyType, PyTypeRef};
6use crate::function::FuncArgs;
7use crate::object::{Traverse, TraverseFn};
8use crate::types::{PyTypeFlags, PyTypeSlots};
9use crate::{
10    AsObject, Context, Py, PyAtomicRef, PyObject, PyObjectRef, PyRef, PyResult, VirtualMachine,
11};
12use core::fmt::Write;
13use rustpython_common::wtf8::Wtf8Buf;
14
15use crate::exceptions::types::PyBaseException;
16
17/// Create dynamic ExceptionGroup type with multiple inheritance
18fn create_exception_group(ctx: &Context) -> PyRef<PyType> {
19    let excs = &ctx.exceptions;
20    let exception_group_slots = PyTypeSlots {
21        flags: crate::types::AtomicPyTypeFlags::from_plain(PyTypeFlags::HEAP_TYPE_WITH_DICT),
22        ..Default::default()
23    };
24    let mut attrs = crate::builtins::type_::PyAttributes::default();
25    attrs.insert(
26        crate::identifier!(ctx, __module__),
27        ctx.intern_str("builtins").to_object(),
28    );
29    PyType::new_heap(
30        "ExceptionGroup",
31        vec![
32            excs.base_exception_group.to_owned(),
33            excs.exception_type.to_owned(),
34        ],
35        attrs,
36        exception_group_slots,
37        ctx.types.type_type.to_owned(),
38        ctx,
39    )
40    .expect("Failed to create ExceptionGroup type with multiple inheritance")
41}
42
43#[must_use]
44pub fn exception_group() -> &'static Py<PyType> {
45    ::rustpython_vm::common::static_cell! {
46        static CELL: ::rustpython_vm::builtins::PyTypeRef;
47    }
48    CELL.get_or_init(|| create_exception_group(Context::genesis()))
49}
50
51pub(super) mod types {
52    use super::*;
53    use crate::PyPayload;
54    use crate::builtins::PyGenericAlias;
55    use crate::types::{Constructor, Initializer};
56
57    #[pyexception(name, base = PyBaseException, ctx = "base_exception_group", traverse = "manual")]
58    #[repr(C)]
59    pub struct PyBaseExceptionGroup {
60        base: PyBaseException,
61        #[pymember(name = "message")]
62        msg: PyAtomicRef<PyObject>,
63        #[pymember(name = "exceptions")]
64        excs: PyAtomicRef<PyObject>,
65        excs_str: PyAtomicRef<Option<PyObject>>,
66    }
67
68    impl crate::class::PySubclass for PyBaseExceptionGroup {
69        type Base = PyBaseException;
70        fn as_base(&self) -> &Self::Base {
71            &self.base
72        }
73    }
74
75    impl core::fmt::Debug for PyBaseExceptionGroup {
76        fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
77            f.debug_struct("PyBaseExceptionGroup")
78                .finish_non_exhaustive()
79        }
80    }
81
82    unsafe impl Traverse for PyBaseExceptionGroup {
83        fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
84            self.base.traverse(tracer_fn);
85            tracer_fn(&self.msg);
86            tracer_fn(&self.excs);
87            if let Some(obj) = self.excs_str.deref() {
88                tracer_fn(obj);
89            }
90        }
91    }
92
93    #[pyexception(with(Constructor, Initializer))]
94    impl PyBaseExceptionGroup {
95        #[pyclassmethod]
96        fn __class_getitem__(
97            cls: PyTypeRef,
98            object: PyObjectRef,
99            vm: &VirtualMachine,
100        ) -> PyResult<PyGenericAlias> {
101            PyGenericAlias::from_args(cls, object, vm)
102        }
103
104        #[pymethod]
105        fn derive(zelf: PyRef<Self>, excs: PyObjectRef, vm: &VirtualMachine) -> PyResult {
106            let message = zelf.msg.to_owned();
107            vm.invoke_exception(vm.ctx.exceptions.base_exception_group, vec![message, excs])
108                .map(|e| e.into())
109        }
110
111        #[pymethod]
112        fn subgroup(
113            zelf: PyRef<Self>,
114            matcher_value: PyObjectRef,
115            vm: &VirtualMachine,
116        ) -> PyResult {
117            let matcher = get_condition_matcher(&matcher_value, vm)?;
118
119            // If self matches the condition entirely, return self
120            let zelf_obj: PyObjectRef = zelf.clone().into();
121            if matcher.check(&zelf_obj, vm)? {
122                return Ok(zelf_obj);
123            }
124
125            let exceptions = get_exceptions_tuple(&zelf, vm)?;
126            let mut matching: Vec<PyObjectRef> = Vec::new();
127            let mut modified = false;
128
129            for exc in exceptions {
130                if is_base_exception_group(&exc, vm) {
131                    // Recursive call for nested groups. It pushes no Python
132                    // frame, so a deep enough group runs off the native stack
133                    // unless this guard is here.
134                    let subgroup_result = vm
135                        .with_recursion("in exception group subgroup", || {
136                            vm.call_method(&exc, "subgroup", (matcher_value.clone(),))
137                        })?;
138                    if !vm.is_none(&subgroup_result) {
139                        matching.push(subgroup_result.clone());
140                    }
141                    if !subgroup_result.is(&exc) {
142                        modified = true;
143                    }
144                } else if matcher.check(&exc, vm)? {
145                    matching.push(exc);
146                } else {
147                    modified = true;
148                }
149            }
150
151            if !modified {
152                return Ok(zelf.into());
153            }
154
155            if matching.is_empty() {
156                return Ok(vm.ctx.none());
157            }
158
159            // Create new group with matching exceptions and copy metadata
160            derive_and_copy_attributes(&zelf, matching, vm)
161        }
162
163        #[pymethod]
164        fn split(
165            zelf: PyRef<Self>,
166            matcher_value: PyObjectRef,
167            vm: &VirtualMachine,
168        ) -> PyResult<PyTupleRef> {
169            let matcher = get_condition_matcher(&matcher_value, vm)?;
170
171            // If self matches the condition entirely
172            let zelf_obj: PyObjectRef = zelf.clone().into();
173            if matcher.check(&zelf_obj, vm)? {
174                return Ok(vm.ctx.new_tuple(vec![zelf_obj, vm.ctx.none()]));
175            }
176
177            let exceptions = get_exceptions_tuple(&zelf, vm)?;
178            let mut matching: Vec<PyObjectRef> = Vec::new();
179            let mut rest: Vec<PyObjectRef> = Vec::new();
180
181            for exc in exceptions {
182                if is_base_exception_group(&exc, vm) {
183                    // Same as in subgroup: nothing else bounds this recursion
184                    // against the native stack.
185                    let result = vm.with_recursion("in exception group split", || {
186                        vm.call_method(&exc, "split", (matcher_value.clone(),))
187                    })?;
188                    let result_tuple: PyTupleRef = result.try_into_value(vm)?;
189                    let match_part = result_tuple
190                        .as_slice()
191                        .first()
192                        .cloned()
193                        .unwrap_or_else(|| vm.ctx.none());
194                    let rest_part = result_tuple
195                        .as_slice()
196                        .get(1)
197                        .cloned()
198                        .unwrap_or_else(|| vm.ctx.none());
199
200                    if !vm.is_none(&match_part) {
201                        matching.push(match_part);
202                    }
203                    if !vm.is_none(&rest_part) {
204                        rest.push(rest_part);
205                    }
206                } else if matcher.check(&exc, vm)? {
207                    matching.push(exc);
208                } else {
209                    rest.push(exc);
210                }
211            }
212
213            let match_group = if matching.is_empty() {
214                vm.ctx.none()
215            } else {
216                derive_and_copy_attributes(&zelf, matching, vm)?
217            };
218
219            let rest_group = if rest.is_empty() {
220                vm.ctx.none()
221            } else {
222                derive_and_copy_attributes(&zelf, rest, vm)?
223            };
224
225            Ok(vm.ctx.new_tuple(vec![match_group, rest_group]))
226        }
227
228        #[pyslot]
229        fn slot_str(zelf: &PyObject, vm: &VirtualMachine) -> PyResult<PyStrRef> {
230            let zelf: &Py<Self> = zelf
231                .downcast_ref()
232                .expect("slot wrapper checked BaseExceptionGroup");
233            let message = zelf.msg.str(vm)?;
234            let num_excs = zelf
235                .excs
236                .downcast_ref::<PyTuple>()
237                .map_or(0, |t| t.as_slice().len());
238
239            let suffix = if num_excs == 1 { "" } else { "s" };
240            let mut result = message.as_wtf8().to_owned();
241            write!(result, " ({num_excs} sub-exception{suffix})")
242                .expect("formatting into a string buffer cannot fail");
243            Ok(vm.ctx.new_str(result))
244        }
245
246        #[pyslot]
247        fn slot_repr(zelf: &PyObject, vm: &VirtualMachine) -> PyResult<PyStrRef> {
248            let zelf = zelf
249                .downcast_ref::<PyBaseExceptionGroup>()
250                .expect("exception group must be BaseExceptionGroup");
251            let class_name = zelf.class().name().to_owned();
252            let message = zelf.msg.repr(vm)?;
253
254            let exceptions_str = if let Some(saved) = zelf.excs_str.load_owned() {
255                saved
256                    .downcast::<crate::builtins::PyStr>()
257                    .map_err(|_| vm.new_type_error("__repr__ returned non-string"))?
258            } else {
259                let args = zelf.base.args();
260                let exceptions_obj = if args.as_slice().len() == 2
261                    && args.as_slice()[1].downcast_ref::<PyList>().is_some()
262                {
263                    let list = match zelf.excs.downcast_ref::<PyTuple>() {
264                        Some(tuple) => vm.ctx.new_list(tuple.as_slice().to_vec()),
265                        None => vm.ctx.new_list(vec![]),
266                    };
267                    list.into()
268                } else {
269                    zelf.excs.to_owned()
270                };
271                exceptions_obj.repr(vm)?
272            };
273
274            let mut result = Wtf8Buf::new();
275            write!(result, "{class_name}(").unwrap();
276            result.push_wtf8(message.as_wtf8());
277            result.push_str(", ");
278            result.push_wtf8(exceptions_str.as_wtf8());
279            result.push_str(")");
280
281            Ok(vm.ctx.new_str(result))
282        }
283    }
284
285    impl Constructor for PyBaseExceptionGroup {
286        type Args = FuncArgs;
287
288        fn slot_new(cls: PyTypeRef, args: FuncArgs, vm: &VirtualMachine) -> PyResult {
289            if args.args.len() != 2 {
290                return Err(vm.new_type_error(format!(
291                    "BaseExceptionGroup.__new__() takes exactly 2 arguments ({} given)",
292                    args.args.len()
293                )));
294            }
295
296            let message = args.args[0].clone();
297            if !message.fast_isinstance(vm.ctx.types.str_type) {
298                return Err(vm.new_type_error(format!(
299                    "argument 1 must be str, not {}",
300                    message.class().name()
301                )));
302            }
303
304            let exceptions_arg = &args.args[1];
305            exceptions_arg.try_sequence(vm).map_err(|_| {
306                vm.new_type_error("second argument (exceptions) must be a sequence")
307            })?;
308
309            let is_list = exceptions_arg.downcast_ref::<PyList>().is_some();
310            let is_tuple = exceptions_arg.downcast_ref::<PyTuple>().is_some();
311            let excs_str = if !is_list && !is_tuple {
312                Some(exceptions_arg.repr(vm)?.into())
313            } else {
314                None
315            };
316
317            let exceptions: Vec<PyObjectRef> = exceptions_arg.try_to_value(vm).map_err(|_| {
318                vm.new_type_error("second argument (exceptions) must be a sequence")
319            })?;
320
321            if exceptions.is_empty() {
322                return Err(
323                    vm.new_value_error("second argument (exceptions) must be a non-empty sequence")
324                );
325            }
326
327            let mut has_non_exception = false;
328            for (i, exc) in exceptions.iter().enumerate() {
329                if !exc.fast_isinstance(vm.ctx.exceptions.base_exception_type) {
330                    return Err(vm.new_value_error(format!(
331                        "Item {i} of second argument (exceptions) is not an exception"
332                    )));
333                }
334                if !exc.fast_isinstance(vm.ctx.exceptions.exception_type) {
335                    has_non_exception = true;
336                }
337            }
338
339            let exception_group_type = crate::exception_group::exception_group();
340
341            let actual_cls = if cls.is(exception_group_type) {
342                if has_non_exception {
343                    return Err(
344                        vm.new_type_error("Cannot nest BaseExceptions in an ExceptionGroup")
345                    );
346                }
347                cls
348            } else if cls.is(vm.ctx.exceptions.base_exception_group) {
349                if !has_non_exception {
350                    exception_group_type.to_owned()
351                } else {
352                    cls
353                }
354            } else {
355                if has_non_exception && cls.fast_issubclass(vm.ctx.exceptions.exception_type) {
356                    return Err(vm.new_type_error(format!(
357                        "Cannot nest BaseExceptions in '{}'",
358                        cls.name()
359                    )));
360                }
361                cls
362            };
363
364            // Keep an exact tuple as-is so `.exceptions is original` for tuples.
365            let exceptions_tuple = if exceptions_arg.class().is(vm.ctx.types.tuple_type) {
366                exceptions_arg
367                    .clone()
368                    .downcast::<PyTuple>()
369                    .expect("exact tuple")
370            } else {
371                vm.ctx.new_tuple(exceptions)
372            };
373
374            let payload = Self {
375                base: PyBaseException::new(args.args.clone(), vm),
376                msg: message.into(),
377                excs: PyObjectRef::from(exceptions_tuple).into(),
378                excs_str: excs_str.into(),
379            };
380            payload
381                .into_ref_with_type_lazy_dict(vm, actual_cls)
382                .map(Into::into)
383        }
384
385        fn py_new(_cls: &Py<PyType>, _args: Self::Args, _vm: &VirtualMachine) -> PyResult<Self> {
386            unimplemented!("use slot_new")
387        }
388    }
389
390    impl Initializer for PyBaseExceptionGroup {
391        type Args = FuncArgs;
392
393        fn slot_init(zelf: &PyObject, args: FuncArgs, vm: &VirtualMachine) -> PyResult<()> {
394            if !args.kwargs.is_empty() {
395                return Err(vm.new_type_error(format!(
396                    "{} does not take keyword arguments",
397                    zelf.class().name()
398                )));
399            }
400            PyBaseException::slot_init(zelf, args, vm)
401        }
402
403        fn init(_zelf: &Py<Self>, _args: Self::Args, _vm: &VirtualMachine) -> PyResult<()> {
404            unreachable!("slot_init is overridden")
405        }
406    }
407
408    // Helper functions for ExceptionGroup
409    fn is_base_exception_group(obj: &PyObject, vm: &VirtualMachine) -> bool {
410        obj.fast_isinstance(vm.ctx.exceptions.base_exception_group)
411    }
412
413    fn get_exceptions_tuple(
414        exc: &Py<PyBaseExceptionGroup>,
415        vm: &VirtualMachine,
416    ) -> PyResult<Vec<PyObjectRef>> {
417        let tuple = exc
418            .excs
419            .downcast_ref::<PyTuple>()
420            .ok_or_else(|| vm.new_type_error("exceptions must be a tuple"))?;
421        Ok(tuple.as_slice().to_vec())
422    }
423
424    enum ConditionMatcher {
425        Type(PyTypeRef),
426        Types(Vec<PyTypeRef>),
427        Callable(PyObjectRef),
428    }
429
430    fn get_condition_matcher(
431        condition: &PyObject,
432        vm: &VirtualMachine,
433    ) -> PyResult<ConditionMatcher> {
434        // If it's a type and subclass of BaseException
435        if let Some(typ) = condition.downcast_ref::<PyType>()
436            && typ.fast_issubclass(vm.ctx.exceptions.base_exception_type)
437        {
438            return Ok(ConditionMatcher::Type(typ.to_owned()));
439        }
440
441        // If it's a tuple of types
442        if let Some(tuple) = condition.downcast_ref::<PyTuple>() {
443            let mut types = Vec::new();
444            for item in tuple {
445                let typ: PyTypeRef = item.clone().try_into_value(vm).map_err(|_| {
446                    vm.new_type_error(
447                        "expected a function, exception type or tuple of exception types",
448                    )
449                })?;
450                if !typ.fast_issubclass(vm.ctx.exceptions.base_exception_type) {
451                    return Err(vm.new_type_error(
452                        "expected a function, exception type or tuple of exception types",
453                    ));
454                }
455                types.push(typ);
456            }
457            if !types.is_empty() {
458                return Ok(ConditionMatcher::Types(types));
459            }
460        }
461
462        // If it's callable (but not a type)
463        if condition.is_callable() && condition.downcast_ref::<PyType>().is_none() {
464            return Ok(ConditionMatcher::Callable(condition.to_owned()));
465        }
466
467        Err(vm.new_type_error("expected a function, exception type or tuple of exception types"))
468    }
469
470    impl ConditionMatcher {
471        fn check(&self, exc: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
472            match self {
473                Self::Type(typ) => Ok(exc.fast_isinstance(typ)),
474                Self::Types(types) => Ok(types.iter().any(|t| exc.fast_isinstance(t))),
475                Self::Callable(func) => {
476                    let result = func.call((exc.to_owned(),), vm)?;
477                    result.try_to_bool(vm)
478                }
479            }
480        }
481    }
482
483    pub(crate) fn derive_and_copy_attributes(
484        orig: &Py<PyBaseExceptionGroup>,
485        excs: Vec<PyObjectRef>,
486        vm: &VirtualMachine,
487    ) -> PyResult<PyObjectRef> {
488        // Call derive method to create new group
489        let excs_seq = vm.ctx.new_list(excs);
490        let new_group = vm.call_method(orig.as_object(), "derive", (excs_seq,))?;
491
492        // Verify derive returned a BaseExceptionGroup
493        if !is_base_exception_group(&new_group, vm) {
494            return Err(vm.new_type_error("derive must return an instance of BaseExceptionGroup"));
495        }
496
497        // Copy traceback
498        if let Some(tb) = orig.base.__traceback__() {
499            new_group.set_attr("__traceback__", tb, vm)?;
500        }
501
502        // Copy context
503        if let Some(ctx) = orig.base.__context__() {
504            new_group.set_attr("__context__", ctx, vm)?;
505        }
506
507        // Copy cause
508        if let Some(cause) = orig.base.__cause__() {
509            new_group.set_attr("__cause__", cause, vm)?;
510        }
511
512        // Copy notes (if present) - make a copy of the list
513        if let Ok(notes) = orig.as_object().get_attr("__notes__", vm)
514            && let Some(notes_list) = notes.downcast_ref::<PyList>()
515        {
516            let notes_copy = vm.ctx.new_list(notes_list.borrow_vec().to_vec());
517            new_group.set_attr("__notes__", notes_copy, vm)?;
518        }
519
520        Ok(new_group)
521    }
522}