diff --git a/mypyc/irbuild/statement.py b/mypyc/irbuild/statement.py index a797b8f5ab34..aac19cc8a75e 100644 --- a/mypyc/irbuild/statement.py +++ b/mypyc/irbuild/statement.py @@ -121,6 +121,7 @@ keep_propagating_op, no_err_occurred_op, propagate_if_error_op, + raise_exception_from_op, raise_exception_op, reraise_exception_op, restore_exc_info_op, @@ -679,7 +680,16 @@ def transform_raise_stmt(builder: IRBuilder, s: RaiseStmt) -> None: return exc = builder.accept(s.expr) - builder.call_c(raise_exception_op, [exc], s.line) + if s.from_expr is not None: + if isinstance(exc, Register): + # Evaluating the cause may reassign the variable holding the exception. + temp = Register(exc.type) + builder.assign(temp, exc, s.line) + exc = temp + cause = builder.accept(s.from_expr) + builder.call_c(raise_exception_from_op, [exc, cause], s.line) + else: + builder.call_c(raise_exception_op, [exc], s.line) builder.add(Unreachable()) diff --git a/mypyc/lib-rt/CPy.h b/mypyc/lib-rt/CPy.h index 5500ebdd86e7..64cd25f0c335 100644 --- a/mypyc/lib-rt/CPy.h +++ b/mypyc/lib-rt/CPy.h @@ -979,6 +979,7 @@ static inline bool CPy_KeepPropagating(void) { #define CPy_ExcState() PyThreadState_GET()->exc_info void CPy_Raise(PyObject *exc); +void CPy_RaiseFrom(PyObject *exc, PyObject *cause); void CPy_Reraise(void); void CPyErr_SetObjectAndTraceback(PyObject *type, PyObject *value, PyObject *traceback); tuple_T3OOO CPy_CatchError(void); diff --git a/mypyc/lib-rt/exc_ops.c b/mypyc/lib-rt/exc_ops.c index c2a220d280ee..f6a79fa8197a 100644 --- a/mypyc/lib-rt/exc_ops.c +++ b/mypyc/lib-rt/exc_ops.c @@ -7,16 +7,67 @@ #include #include "CPy.h" -void CPy_Raise(PyObject *exc) { - if (PyObject_IsInstance(exc, (PyObject *)&PyType_Type)) { - PyObject *obj = PyObject_CallNoArgs(exc); - if (!obj) +// Return a new exception instance, or NULL with an error set. +static PyObject *instantiate_exception(PyObject *type) { + PyObject *value = PyObject_CallNoArgs(type); + if (!value) + return NULL; + if (!PyExceptionInstance_Check(value)) { + PyErr_Format(PyExc_TypeError, + "calling %R should have returned an instance of " + "BaseException, not %R", type, Py_TYPE(value)); + Py_DECREF(value); + return NULL; + } + return value; +} + +// A NULL cause means no 'from' clause; Py_None means 'from None'. +static void raise_exception(PyObject *exc, PyObject *cause) { + PyObject *type; + PyObject *value; + if (PyExceptionClass_Check(exc)) { + type = exc; + value = instantiate_exception(exc); + if (!value) return; - PyErr_SetObject(exc, obj); - Py_DECREF(obj); + } else if (PyExceptionInstance_Check(exc)) { + type = (PyObject *)Py_TYPE(exc); + value = Py_NewRef(exc); } else { - PyErr_SetObject((PyObject *)Py_TYPE(exc), exc); + PyErr_SetString(PyExc_TypeError, "exceptions must derive from BaseException"); + return; + } + + if (cause != NULL) { + PyObject *fixed_cause; + if (PyExceptionClass_Check(cause)) { + fixed_cause = instantiate_exception(cause); + if (!fixed_cause) + goto fail; + } else if (PyExceptionInstance_Check(cause)) { + fixed_cause = Py_NewRef(cause); + } else if (Py_IsNone(cause)) { + fixed_cause = NULL; + } else { + PyErr_SetString(PyExc_TypeError, "exception causes must derive from BaseException"); + goto fail; + } + // This steals the reference to the cause and sets __suppress_context__ + PyException_SetCause(value, fixed_cause); } + // Normalize against the original class, even if __new__ returned an unrelated exception. + PyErr_SetObject(type, value); +fail: + Py_DECREF(value); +} + +void CPy_Raise(PyObject *exc) { + raise_exception(exc, NULL); +} + +void CPy_RaiseFrom(PyObject *exc, PyObject *cause) { + raise_exception(exc, cause); } void CPy_Reraise(void) { diff --git a/mypyc/primitives/exc_ops.py b/mypyc/primitives/exc_ops.py index e1234f807afa..528a89ef2331 100644 --- a/mypyc/primitives/exc_ops.py +++ b/mypyc/primitives/exc_ops.py @@ -15,6 +15,14 @@ error_kind=ERR_ALWAYS, ) +# Like raise_exception_op, but also set the cause (raise from ). +raise_exception_from_op = custom_op( + arg_types=[object_rprimitive, object_rprimitive], + return_type=void_rtype, + c_function_name="CPy_RaiseFrom", + error_kind=ERR_ALWAYS, +) + # Raise StopIteration exception with the specified value (which can be NULL). set_stop_iteration_value = custom_op( arg_types=[object_rprimitive], diff --git a/mypyc/test-data/fixtures/ir.py b/mypyc/test-data/fixtures/ir.py index edb9318e2e39..708e9356d203 100644 --- a/mypyc/test-data/fixtures/ir.py +++ b/mypyc/test-data/fixtures/ir.py @@ -351,7 +351,10 @@ def fget(self) -> Any: ... def fset(self, value: Any) -> None: ... def fdel(self) -> None: ... -class BaseException: pass +class BaseException: + __cause__: Optional[BaseException] + __context__: Optional[BaseException] + __suppress_context__: bool class Exception(BaseException): def __init__(self, message: Optional[str] = None) -> None: pass diff --git a/mypyc/test-data/run-exceptions.test b/mypyc/test-data/run-exceptions.test index 1b180b933197..8257bfd9b894 100644 --- a/mypyc/test-data/run-exceptions.test +++ b/mypyc/test-data/run-exceptions.test @@ -532,3 +532,136 @@ Traceback (most recent call last): File "native.py", line 6, in f(y) TypeError: int object expected; got str + +[case testRaiseFrom] +from typing import Any + +from exceptions_helper import OtherException +from testutil import assertRaises + +class Inner(Exception): + pass + +class Outer(Exception): + pass + +def raise_from(exc: Any, cause: Any) -> None: + try: + raise Inner("inner") + except Inner: + raise exc from cause + +def test_raise_from_instance() -> None: + cause = ValueError("cause") + try: + raise_from(Outer("outer"), cause) + except Outer as e: + assert e.__cause__ is cause + assert e.__suppress_context__ + assert isinstance(e.__context__, Inner) + else: + assert False + +def test_raise_from_class() -> None: + try: + raise_from(Outer, ValueError) + except Outer as e: + assert type(e.__cause__) is ValueError + assert e.__suppress_context__ + else: + assert False + +def test_raise_from_none() -> None: + try: + raise_from(Outer("outer"), None) + except Outer as e: + assert e.__cause__ is None + assert e.__suppress_context__ + assert isinstance(e.__context__, Inner) + else: + assert False + +def test_raise_from_invalid_cause() -> None: + with assertRaises(TypeError, "exception causes must derive from BaseException"): + raise_from(Outer("outer"), 1) + +def test_raise_from_invalid_exception() -> None: + with assertRaises(TypeError, "exceptions must derive from BaseException"): + raise_from(1, ValueError()) + +def test_raise_from_in_except() -> None: + try: + try: + raise Inner("inner") + except Inner as exc: + raise Outer("outer") from exc + except Outer as e: + assert isinstance(e.__cause__, Inner) + assert e.__cause__ is e.__context__ + else: + assert False + +def test_raise_from_reassigned_exception() -> None: + original = ValueError("original") + cause = TypeError("cause") + exc: Exception = original + try: + raise exc from (exc := cause) + except ValueError as e: + assert e is original + assert e.__cause__ is cause + assert exc is cause + else: + assert False + +def test_raise_from_class_returning_unrelated_exception() -> None: + # Normalization calls OtherException again with the ValueError as an argument. + with assertRaises(TypeError): + raise_from(OtherException, None) + +def raise_plain(exc: Any) -> None: + raise exc + +def test_raise_plain_preserves_cause() -> None: + cause = ValueError("cause") + exc = Outer("outer") + try: + raise_from(exc, cause) + except Outer: + pass + try: + raise_plain(exc) + except Outer as e: + assert e is exc + assert e.__cause__ is cause + assert e.__suppress_context__ + else: + assert False + +def test_raise_plain_does_not_suppress_context() -> None: + try: + try: + raise Inner("inner") + except Inner: + raise_plain(Outer("outer")) + except Outer as e: + assert e.__cause__ is None + assert not e.__suppress_context__ + assert isinstance(e.__context__, Inner) + else: + assert False + +def test_raise_plain_invalid_exception() -> None: + with assertRaises(TypeError, "exceptions must derive from BaseException"): + raise_plain(1) + +def test_raise_plain_class_returning_unrelated_exception() -> None: + with assertRaises(TypeError): + raise_plain(OtherException) + +[file exceptions_helper.py] +from typing import Any + +class OtherException(Exception): + def __new__(cls) -> Any: + return ValueError("other")