1use 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
17fn 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 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 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 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 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 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 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 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 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 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 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 let excs_seq = vm.ctx.new_list(excs);
490 let new_group = vm.call_method(orig.as_object(), "derive", (excs_seq,))?;
491
492 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 if let Some(tb) = orig.base.__traceback__() {
499 new_group.set_attr("__traceback__", tb, vm)?;
500 }
501
502 if let Some(ctx) = orig.base.__context__() {
504 new_group.set_attr("__context__", ctx, vm)?;
505 }
506
507 if let Some(cause) = orig.base.__cause__() {
509 new_group.set_attr("__cause__", cause, vm)?;
510 }
511
512 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}