]> git.ipfire.org Git - thirdparty/Python/cpython.git/commitdiff
GH-148874: Make sure that mngr.__exit__() is always called in a with statement (GH...
authorMark Shannon <Mark.Shannon@arm.com>
Wed, 1 Jul 2026 14:15:33 +0000 (15:15 +0100)
committerGitHub <noreply@github.com>
Wed, 1 Jul 2026 14:15:33 +0000 (15:15 +0100)
* Even if there is an interrupt during the call to mngr.__enter__()

Lib/test/test_with.py
Misc/NEWS.d/next/Core_and_Builtins/2026-06-04-12-53-10.gh-issue-148874.r121cG.rst [new file with mode: 0644]
Modules/_testinternalcapi.c
Modules/_testinternalcapi/test_cases.c.h
Python/bytecodes.c
Python/ceval_macros.h
Python/generated_cases.c.h
Python/optimizer.c
Tools/c-analyzer/cpython/ignored.tsv

index f16611b29a2658c8a2c4fabccc7a2632b0a381ef..60aaa5cd548cbf02ced6320688184ae49c329166 100644 (file)
@@ -10,6 +10,7 @@ import traceback
 import unittest
 from collections import deque
 from contextlib import _GeneratorContextManager, contextmanager, nullcontext
+from _testinternalcapi import SelfInterruptingContextManager
 
 
 def do_with(obj):
@@ -850,5 +851,21 @@ class NestedWith(unittest.TestCase):
                                  expected)
 
 
+class InterruptDuringEnter(unittest.TestCase):
+
+    def test_exit_called_after_interrupt(self):
+        cm = SelfInterruptingContextManager()
+        self.assertFalse(cm.within())
+        try:
+            with cm:
+                self.assertTrue(cm.within())
+        except KeyboardInterrupt:
+            self.assertFalse(cm.within())
+            return
+        except:
+            self.fail("Wrong exception raised")
+        self.fail("No exception raised")
+
+
 if __name__ == '__main__':
     unittest.main()
diff --git a/Misc/NEWS.d/next/Core_and_Builtins/2026-06-04-12-53-10.gh-issue-148874.r121cG.rst b/Misc/NEWS.d/next/Core_and_Builtins/2026-06-04-12-53-10.gh-issue-148874.r121cG.rst
new file mode 100644 (file)
index 0000000..95f9333
--- /dev/null
@@ -0,0 +1,3 @@
+Ignore interrupts immediately after calling the ``__enter__`` method of a
+context menager in a ``with`` statement. This ensures that the ``__exit__``
+method is always called in a ``with`` statement.
index f6ff7820821ce128a5e2a1d5c8e554777c9371ed..ea3ad2b81c28668d9d345acf098e9231566ae114 100644 (file)
@@ -3207,6 +3207,66 @@ test_thread_state_ensure_from_view_interp_switch(PyObject *self, PyObject *unuse
     Py_RETURN_NONE;
 }
 
+/* Self interrupting context manager */
+
+typedef struct {
+    PyObject_HEAD
+    int within;
+} SelfInterruptingContextManagerObject;
+
+static PyObject *
+new_self_interrupting(PyTypeObject *type, PyObject *args, PyObject *kwds)
+{
+    SelfInterruptingContextManagerObject *self =
+        (SelfInterruptingContextManagerObject *)type->tp_alloc(type, 0);
+    if (self != NULL) {
+        self->within = 0;
+    }
+    return (PyObject *)self;
+}
+
+static PyObject *
+self_interrupting_enter(PyObject *op, PyObject *Py_UNUSED(dummy))
+{
+    ((SelfInterruptingContextManagerObject *)op)->within = 1;
+    PyThreadState *tstate = PyThreadState_Get();
+    PyObject *ki = Py_NewRef(PyExc_KeyboardInterrupt);
+    PyObject *old_exc = _Py_atomic_exchange_ptr(&tstate->async_exc, ki);
+    _Py_set_eval_breaker_bit(tstate, _PY_ASYNC_EXCEPTION_BIT);
+    Py_XDECREF(old_exc);
+
+    return Py_NewRef(op);
+}
+
+static PyObject *
+self_interrupting_within(PyObject *op, PyObject *Py_UNUSED(dummy))
+{
+    return PyBool_FromLong(((SelfInterruptingContextManagerObject *)op)->within);
+}
+
+static PyObject *
+self_interrupting_exit(PyObject *op, PyObject *Py_UNUSED(args)) {
+    ((SelfInterruptingContextManagerObject *)op)->within = 0;
+    Py_RETURN_NONE;
+}
+
+static PyMethodDef self_interrupting_methods[] = {
+    {"__enter__", self_interrupting_enter, METH_NOARGS, NULL},
+    {"within", self_interrupting_within, METH_NOARGS, NULL},
+    {"__exit__", self_interrupting_exit, METH_VARARGS, NULL},
+    {NULL, NULL} /* sentinel */
+};
+
+static PyTypeObject SelfInterruptingContextManager_Type = {
+    PyVarObject_HEAD_INIT(NULL, 0)
+    "_testcapi.SelfInterruptingContextManager",
+    sizeof(SelfInterruptingContextManagerObject),
+    .tp_flags = Py_TPFLAGS_DEFAULT | Py_TPFLAGS_IMMUTABLETYPE,
+    .tp_new = new_self_interrupting,
+    .tp_methods = self_interrupting_methods,
+};
+
+
 static PyMethodDef module_functions[] = {
     {"get_configs", get_configs, METH_NOARGS},
     {"get_eval_frame_stats", get_eval_frame_stats, METH_NOARGS, NULL},
@@ -3429,6 +3489,11 @@ module_exec(PyObject *module)
     }
 #endif
 
+    if (PyType_Ready(&SelfInterruptingContextManager_Type) < 0) {
+        return 1;
+    }
+    PyModule_AddObject(module, "SelfInterruptingContextManager", (PyObject *)&SelfInterruptingContextManager_Type);
+
     return 0;
 }
 
index f36c8192ff26623af5a1139fbd3d1a22de352e7e..5f2b1ae5d978aaba5cc039b777db22490600fb17 100644 (file)
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
index 753530b0dabff319717969cead598cf79f5c8dc6..8ac397081e27384f60fee780bdd6e78bafd8030c 100644 (file)
@@ -161,7 +161,7 @@ dummy_func(
         }
 
         replaced op(_CHECK_PERIODIC_AT_END, (--)) {
-            int err = check_periodics(tstate);
+            int err = check_periodics_at_end(tstate, frame);
             ERROR_IF(err != 0);
         }
 
index b13884bf8214d413d84b75f05c5862f4de33e70e..f19adfa0cfcfc1523a2e06f38e64e691d8f8ff88 100644 (file)
@@ -528,6 +528,22 @@ check_periodics(PyThreadState *tstate) {
     return 0;
 }
 
+static inline int
+check_periodics_at_end(PyThreadState *tstate, _PyInterpreterFrame *frame) {
+    _Py_CHECK_EMSCRIPTEN_SIGNALS_PERIODICALLY();
+    QSBR_QUIESCENT_STATE(tstate);
+    if (_Py_atomic_load_uintptr_relaxed(&tstate->eval_breaker) & _PY_EVAL_EVENTS_MASK) {
+        // Do not handle pending interrupts if the previous instruction was LOAD_SPECIAL
+        // This may also not handle interrupts if a cache looks like LOAD_SPECIAL,
+        // but this is benign as we won't skip periodic checks indefinitely.
+        if (frame->instr_ptr[-1].op.code == LOAD_SPECIAL) {
+            return 0;
+        }
+        return _Py_HandlePending(tstate);
+    }
+    return 0;
+}
+
 // Mark the generator as executing. Returns true if the state was changed,
 // false if it was already executing or finished.
 static inline bool
index 88678f14a99585f249899fc5a29941043bcf12bd..cc95179ccaab17cbdf868a2d54256f9f46c8a888 100644 (file)
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
             {
                 assert(stack_pointer == _PyFrame_GetStackPointer(frame));
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
                 ASSERT_WITHIN_STACK_BOUNDS(__FILE__, __LINE__);
                 _PyFrame_SetStackPointer(frame, stack_pointer);
                 _PyFrame_StackPointerValidate(frame);
-                int err = check_periodics(tstate);
+                int err = check_periodics_at_end(tstate, frame);
                 _PyFrame_StackPointerInvalidate(frame);
                 if (err != 0) {
                     JUMP_TO_LABEL(error);
index e95e4b5e24b2c54fc8d07130652b1c6da36e67bc..c9f6ebdb62f07b203218dbf9bd062a37128fa22b 100644 (file)
@@ -956,10 +956,15 @@ _PyJit_translate_single_bytecode_to_trace(
                     case OPARG_REPLACED:
                         uop = _PyUOp_Replacements[uop];
                         assert(uop != 0);
-
                         uint32_t next_inst = target + 1 + _PyOpcode_Caches[_PyOpcode_Deopt[opcode]];
                         if (uop == _TIER2_RESUME_CHECK) {
-                            target = next_inst;
+                            if (this_instr[-1].op.code == LOAD_SPECIAL) {
+                                // Don't check eval breaker immediately after LOAD_SPECIAL
+                                uop = _NOP;
+                            }
+                            else {
+                                target = next_inst;
+                            }
                         }
                         else {
                             int extended_arg = orig_oparg > 255;
index bf08e5568205e7a03471e4aee75062d3b923b69a..6e18593ad698570de4a2581de08c931eb1751790 100644 (file)
@@ -577,6 +577,7 @@ Modules/_testimportmultiple.c       -       _testimportmultiple     -
 Modules/_testinternalcapi.c    -       pending_identify_result -
 Modules/_testinternalcapi.c    -       Test_EvalFrame_Resumes  -
 Modules/_testinternalcapi.c    -       Test_EvalFrame_Loads    -
+Modules/_testinternalcapi.c    -       SelfInterruptingContextManager_Type     -
 Modules/_testinternalcapi/interpreter.c        -       Test_EvalFrame_Resumes  -
 Modules/_testinternalcapi/interpreter.c        -       Test_EvalFrame_Loads    -
 Modules/_testlimitedcapi/slots.c       -       TestMethods     -