Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion Lib/test/test_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -1882,7 +1882,6 @@ class D(dict):
D.__getitem__ = dict.__getitem__
self.assertIs(d[None], None)

@unittest.expectedFailure # TODO: RUSTPYTHON; AssertionError: <class 'tuple'> != <class 'test.test_types.ClassCreationTests.test_tu[41 chars]ass'>
def test_tuple_subclass_as_bases(self):
# gh-132176: it used to crash on using
# tuple subclass for as base classes.
Expand Down
2 changes: 1 addition & 1 deletion crates/vm/src/builtins/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,7 @@ pub(crate) mod super_;
pub use super_::PySuper;
#[path = "type.rs"]
pub(crate) mod type_;
pub use type_::{PyType, PyTypeRef};
pub use type_::{PyType, PyTypeBases, PyTypeRef};
pub(crate) mod range;
pub use range::PyRange;
pub(crate) mod set;
Expand Down
139 changes: 99 additions & 40 deletions crates/vm/src/builtins/type.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,11 +42,54 @@ use num_traits::ToPrimitive;
use rustpython_common::wtf8::Wtf8;
use std::collections::HashSet;

type PyTypeTupleRef = PyRef<PyTuple<PyTypeRef>>;

#[derive(Clone, Debug)]
pub enum PyTypeBases {
/// Bases created before a Python tuple can be allocated.
Bootstrap(Vec<PyTypeRef>),
/// The Python-visible tuple, with every element validated as a type.
Tuple(PyTypeTupleRef),
}

impl Default for PyTypeBases {
fn default() -> Self {
Self::Bootstrap(Vec::new())
}
}

impl Deref for PyTypeBases {
type Target = [PyTypeRef];

fn deref(&self) -> &Self::Target {
match self {
Self::Bootstrap(bases) => bases,
Self::Tuple(bases) => bases.as_slice(),
}
}
}

unsafe impl Traverse for PyTypeBases {
fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
match self {
Self::Bootstrap(bases) => bases.traverse(tracer_fn),
Self::Tuple(bases) => tracer_fn(bases.as_untyped().as_object()),
}
}

fn clear(&mut self, out: &mut Vec<PyObjectRef>) {
match core::mem::take(self) {
Self::Bootstrap(bases) => out.extend(bases.into_iter().map(Into::into)),
Self::Tuple(bases) => out.push(bases.into_untyped().into()),
}
}
}

#[pyclass(module = false, name = "type", traverse = "manual")]
pub struct PyType {
/// tp_base. Written under the type lock (see `set_bases`); read lock-free.
pub base: PyAtomicRef<Option<Self>>,
pub bases: PyRwLock<Vec<PyTypeRef>>,
pub bases: PyRwLock<PyTypeBases>,
pub mro: PyRwLock<Vec<PyTypeRef>>,
pub subclasses: PyRwLock<Vec<PyRef<PyWeak>>>,
pub attributes: PyRwLock<PyAttributes>,
Expand Down Expand Up @@ -261,10 +304,8 @@ unsafe impl crate::object::Traverse for PyType {
if let Some(base) = unsafe { self.base.swap(None) } {
out.push(base.into());
}
if let Some(mut guard) = self.bases.try_write() {
for base in guard.drain(..) {
out.push(base.into());
}
if let Some(mut bases) = self.bases.try_write() {
bases.clear(out);
}
if let Some(mut guard) = self.mro.try_write() {
for typ in guard.drain(..) {
Expand Down Expand Up @@ -590,6 +631,7 @@ impl PyType {
type_data: PyRwLock::new(None),
specialization_cache: TypeSpecializationCache::new(),
};
let bases = PyTuple::new_ref_typed(bases, ctx);
let base = bases[0].clone();

Self::new_heap_inner(base, bases, attrs, slots, heaptype_ext, metaclass, ctx)
Expand Down Expand Up @@ -758,7 +800,7 @@ impl PyType {
#[allow(clippy::too_many_arguments)]
fn new_heap_inner(
base: PyRef<Self>,
bases: Vec<PyRef<Self>>,
bases: PyTypeTupleRef,
attrs: PyAttributes,
mut slots: PyTypeSlots,
heaptype_ext: HeapTypeExt,
Expand Down Expand Up @@ -814,7 +856,7 @@ impl PyType {
let new_type = PyRef::new_ref(
Self {
base: Some(base).into(),
bases: PyRwLock::new(bases),
bases: PyRwLock::new(PyTypeBases::Tuple(bases)),
mro: PyRwLock::new(mro),
subclasses: PyRwLock::default(),
attributes: PyRwLock::new(attrs),
Expand Down Expand Up @@ -872,7 +914,7 @@ impl PyType {
}

let inherited_abc_tpflags = Self::inherited_abc_tpflags(core::slice::from_ref(&base));
let bases = PyRwLock::new(vec![base.clone()]);
let bases = PyRwLock::new(PyTypeBases::Bootstrap(vec![base.clone()]));
let mro = base.mro_map_collect(|x| x.to_owned());

let new_type = PyRef::new_ref(
Expand Down Expand Up @@ -1466,16 +1508,24 @@ impl Py<PyType> {
impl PyType {
#[pygetset]
fn __bases__(&self, vm: &VirtualMachine) -> PyTupleRef {
vm.ctx.new_tuple(
self.bases
.read()
.iter()
.map(|x| x.as_object().to_owned())
.collect(),
)
let bases = Self::with_type_lock(vm, || self.bases.read().clone());
let types = match bases {
PyTypeBases::Tuple(tuple) => return tuple.into_untyped(),
PyTypeBases::Bootstrap(types) => types,
};

let tuple = PyTuple::new_ref_typed(types, &vm.ctx);
Self::with_type_lock(vm, || {
let mut bases = self.bases.write();
if let PyTypeBases::Tuple(current) = &*bases {
return current.clone().into_untyped();
}
*bases = PyTypeBases::Tuple(tuple.clone());
tuple.into_untyped()
})
}
#[pygetset(setter, name = "__bases__")]
fn set_bases(zelf: &Py<Self>, bases: Vec<PyTypeRef>, vm: &VirtualMachine) -> PyResult<()> {
fn set_bases(zelf: &Py<Self>, bases_tuple: PyTupleRef, vm: &VirtualMachine) -> PyResult<()> {
// TODO: Assigning to __bases__ is only used in typing.NamedTupleMeta.__new__
// Rather than correctly re-initializing the class, we are skipping a few steps for now
if zelf.slots.flags.has_feature(PyTypeFlags::IMMUTABLETYPE) {
Expand All @@ -1484,12 +1534,22 @@ impl PyType {
zelf.name()
)));
}
if bases.is_empty() {
if bases_tuple.is_empty() {
return Err(vm.new_type_error(format!(
"can only assign non-empty tuple to {}.__bases__, not ()",
zelf.name()
)));
}
for base in bases_tuple.iter() {
if base.downcast_ref::<Self>().is_none() {
return Err(vm.new_type_error(format!(
"{}.__bases__ must be tuple of classes, not '{}'",
zelf.name(),
base.class().name()
)));
}
}
let bases = bases_tuple.try_into_typed::<Self>(vm)?;

// TODO: check for mro cycles

Expand Down Expand Up @@ -1554,7 +1614,8 @@ impl PyType {
*subclasses = kept;
}

let old_bases = core::mem::replace(&mut *zelf.bases.write(), bases);
let mut old_bases =
core::mem::replace(&mut *zelf.bases.write(), PyTypeBases::Tuple(bases));
let old_base = unsafe { zelf.base.swap(Some(new_base)) };

// Recursively update the mros of this class and all subclasses,
Expand Down Expand Up @@ -1590,12 +1651,12 @@ impl PyType {
retired.extend(failed_mro.into_iter().map(Into::into));
retired.push(cls.into());
}
let failed_bases = core::mem::replace(&mut *zelf.bases.write(), old_bases);
let mut failed_bases = core::mem::replace(&mut *zelf.bases.write(), old_bases);
if let Some(failed_base) = unsafe { zelf.base.swap(old_base) } {
keep_alive(failed_base, &mut retired);
}
register_subclasses(&zelf.bases.read());
retired.extend(failed_bases.into_iter().map(Into::into));
failed_bases.clear(&mut retired);
zelf.modified_inner();
return Err(err);
}
Expand All @@ -1605,7 +1666,7 @@ impl PyType {
retired.extend(old_mro.into_iter().map(Into::into));
retired.push(cls.into());
}
retired.extend(old_bases.into_iter().map(Into::into));
old_bases.clear(&mut retired);
if let Some(old_base) = old_base {
keep_alive(old_base, &mut retired);
}
Expand Down Expand Up @@ -2107,26 +2168,24 @@ impl Constructor for PyType {

let (metatype, base, bases, base_is_type) = if bases.is_empty() {
let base = vm.ctx.types.object_type.to_owned();
(metatype, base.clone(), vec![base], false)
let bases = PyTuple::new_ref_typed(vec![base.clone()], &vm.ctx);
(metatype, base, bases, false)
} else {
let bases = bases
.iter()
.map(|obj| {
obj.clone().downcast::<Self>().or_else(|obj| {
if vm
.get_attribute_opt(obj, identifier!(vm, __mro_entries__))?
.is_some()
{
Err(vm.new_type_error(
"type() doesn't support MRO entry resolution; \
use types.new_class()",
))
} else {
Err(vm.new_type_error("bases must be types"))
}
})
})
.collect::<PyResult<Vec<_>>>()?;
for obj in bases.iter() {
if obj.downcast_ref::<Self>().is_none() {
if vm
.get_attribute_opt(obj.clone(), identifier!(vm, __mro_entries__))?
.is_some()
{
return Err(vm.new_type_error(
"type() doesn't support MRO entry resolution; \
use types.new_class()",
));
}
return Err(vm.new_type_error("bases must be types"));
}
}
let bases = bases.try_into_typed::<Self>(vm)?;

// Search the bases for the proper metatype to deal with this:
let winner = calculate_meta_class(metatype.clone(), &bases, vm)?;
Expand Down
7 changes: 4 additions & 3 deletions crates/vm/src/object/core.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ use super::{
};
use crate::object::traverse_object::PyObjVTable;
use crate::{
builtins::{PyDictRef, PyType, PyTypeRef},
builtins::{PyDictRef, PyType, PyTypeBases, PyTypeRef},
common::{
atomic::{Ordering, PyAtomic, Radium},
linked_list::{Link, Pointers},
Expand Down Expand Up @@ -2681,7 +2681,8 @@ pub(crate) fn init_type_hierarchy() -> (PyTypeRef, PyTypeRef, PyTypeRef) {
// object's mro is [object]
(*object_type_ptr).payload.mro = PyRwLock::new(vec![object_type.clone()]);

(*type_type_ptr).payload.bases = PyRwLock::new(vec![object_type.clone()]);
(*type_type_ptr).payload.bases =
PyRwLock::new(PyTypeBases::Bootstrap(vec![object_type.clone()]));
(*type_type_ptr).payload.base = Some(object_type.clone()).into();

let type_type = PyTypeRef::from_raw(type_type_ptr.cast());
Expand All @@ -2695,7 +2696,7 @@ pub(crate) fn init_type_hierarchy() -> (PyTypeRef, PyTypeRef, PyTypeRef) {

let weakref_type = PyType {
base: Some(object_type.clone()).into(),
bases: PyRwLock::new(vec![object_type.clone()]),
bases: PyRwLock::new(PyTypeBases::Bootstrap(vec![object_type.clone()])),
mro: PyRwLock::new(vec![object_type.clone()]),
subclasses: PyRwLock::default(),
attributes: PyRwLock::default(),
Expand Down
Loading