From 71f06c4ff1c80e015c3da7b97da439815ca0c80f Mon Sep 17 00:00:00 2001 From: blueloveTH Date: Wed, 29 Jul 2026 19:16:24 +0800 Subject: [PATCH] fix a bug of `__init__` --- src/interpreter/vm.c | 19 +++++++++++-------- tests/410_class_ex.py | 12 +++++++++++- 2 files changed, 22 insertions(+), 9 deletions(-) diff --git a/src/interpreter/vm.c b/src/interpreter/vm.c index 47578ded..dd4d7acb 100644 --- a/src/interpreter/vm.c +++ b/src/interpreter/vm.c @@ -579,8 +579,9 @@ FrameResult VM__vectorcall(VM* self, uint16_t argc, uint16_t kwargc, bool opcall } if(p0->type == tp_type) { + py_Type p0_type = py_totype(p0); // [cls, NULL, args..., kwargs...] - py_Ref new_f = py_tpfindmagic(py_totype(p0), __new__); + py_Ref new_f = py_tpfindmagic(p0_type, __new__); assert(new_f && py_isnil(p0 + 1)); bool is_default_new = new_f->type == tp_nativefunc && new_f->_cfunc == pk__object_new; @@ -598,14 +599,16 @@ FrameResult VM__vectorcall(VM* self, uint16_t argc, uint16_t kwargc, bool opcall // NOTE: previously we use `get_unbound_method` but here we just use `tpfindmagic` // >> [cls, NULL, args..., kwargs...] // >> py_retval() is the new instance - py_Ref init_f = py_tpfindmagic(py_totype(p0), __init__); + py_Ref init_f = py_tpfindmagic(p0_type, __init__); if(init_f) { - // do an inplace patch - *p0 = *init_f; // __init__ - p0[1] = self->last_retval; // self - // [__init__, self, args..., kwargs...] - if(VM__vectorcall(self, argc, kwargc, false) == RES_ERROR) return RES_ERROR; - *py_retval() = p0[1]; // restore the new instance + if(py_isinstance(py_retval(), p0_type)) { + // do an inplace patch + *p0 = *init_f; // __init__ + p0[1] = self->last_retval; // self + // [__init__, self, args..., kwargs...] + if(VM__vectorcall(self, argc, kwargc, false) == RES_ERROR) return RES_ERROR; + *py_retval() = p0[1]; // restore the new instance + } } else { if(is_default_new) { if(argc != 0 || kwargc != 0) { diff --git a/tests/410_class_ex.py b/tests/410_class_ex.py index e52d811d..13d4b643 100644 --- a/tests/410_class_ex.py +++ b/tests/410_class_ex.py @@ -167,4 +167,14 @@ class DerivedClass(BaseClass): return super().f() -assert DerivedClass.f() == 'BaseClass' \ No newline at end of file +assert DerivedClass.f() == 'BaseClass' + +# bad __init__ +class A: + def __new__(cls, *args, **kwargs): + return 1 + + def __init__(self): + assert False + +A()