mod field;
use std::{fmt::Write, mem};
pub(crate) use self::field::DataclassField;
use crate::{
args::{ArgValues, KwargsValues},
bytecode::{CallResult, VM},
defer_drop, defer_drop_mut,
exception_private::{ExcType, ExcTypeExt, RunError, RunResult, SimpleException},
heap::{DropGuard, DropWithContext, HeapData, HeapId, HeapRead, HeapReadOutput},
intern::{StaticStrings, StringId},
modules::ModuleFunctions,
types::{
Class, Dict, Instance, LazyHeapSet, Module,
dataclass::write_dataclass_repr,
instance::{class_name, instance_attr},
},
value::Value,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, strum::Display, serde::Serialize, serde::Deserialize)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum DataclassesFunctions {
Dataclass,
IsDataclass,
}
pub fn create_module(vm: &mut VM<'_>) -> HeapId {
let mut module = Module::new(StaticStrings::Dataclasses);
module.set_attr(
StaticStrings::Dataclass,
Value::ModuleFunction(ModuleFunctions::Dataclasses(DataclassesFunctions::Dataclass)),
vm,
);
module.set_attr(
StaticStrings::IsDataclass,
Value::ModuleFunction(ModuleFunctions::Dataclasses(DataclassesFunctions::IsDataclass)),
vm,
);
vm.heap.allocate(HeapData::Module(Box::new(module)))
}
pub(super) fn call(vm: &mut VM<'_>, func: DataclassesFunctions, args: ArgValues) -> RunResult<Value> {
match func {
DataclassesFunctions::Dataclass => dataclass_decorator(vm, args),
DataclassesFunctions::IsDataclass => is_dataclass(vm, args),
}
}
fn dataclass_decorator(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
if matches!(args, ArgValues::Kwargs(_) | ArgValues::ArgsKargs { .. }) {
args.drop_with(vm);
return Err(ExcType::not_implemented(
"dataclass() keyword options (eq, order, frozen, unsafe_hash, ...) are not yet supported",
)
.into());
}
let cls = args.get_one_arg("dataclass", vm.heap)?;
let mut guard = DropGuard::new(cls, vm);
let (cls, vm) = guard.as_parts();
let Value::Ref(class_id) = cls else {
return Err(non_class_error(cls, vm));
};
let HeapReadOutput::Class(mut class) = vm.heap.read(*class_id) else {
return Err(non_class_error(cls, vm));
};
let fields = build_dataclass_fields(&class, vm)?;
store_dataclass_fields(&mut class, fields, vm)?;
Ok(guard.into_inner())
}
fn non_class_error(cls: &Value, vm: &VM<'_>) -> RunError {
let type_name = cls.py_type_name(vm);
SimpleException::new_msg(
ExcType::TypeError,
format!("dataclass() should be called on a class, not '{type_name}'"),
)
.into()
}
fn build_dataclass_fields<'h>(class: &HeapRead<'h, Class>, vm: &mut VM<'h>) -> RunResult<Value> {
let fields = collect_annotated_fields(class, vm);
let mut guard = DropGuard::new(fields, vm);
let (fields, vm) = guard.as_parts();
validate_fields(vm, fields)?;
reject_unsupported_members(class, vm)?;
let (fields, vm) = guard.into_parts();
allocate_fields_dict(vm, fields)
}
fn allocate_fields_dict(vm: &mut VM<'_>, fields: Vec<DataclassField>) -> RunResult<Value> {
let dict_id = vm.heap.allocate(HeapData::Dict(Dict::with_capacity(fields.len())));
let mut guard = DropGuard::new(Value::Ref(dict_id), vm);
let vm = guard.ctx();
for field in fields {
let name = field.name();
let field_id = vm.heap.allocate(HeapData::DataclassField(field));
let HeapReadOutput::Dict(mut dict) = vm.heap.read(dict_id) else {
unreachable!("the dict was just allocated")
};
let replaced = dict.set(Value::InternString(name), Value::Ref(field_id), vm)?;
replaced.drop_with(vm);
}
Ok(guard.into_inner())
}
fn store_dataclass_fields<'h>(class: &mut HeapRead<'h, Class>, fields: Value, vm: &mut VM<'h>) -> RunResult<()> {
let replaced = class.set_attr(StaticStrings::DataclassFields.into(), fields, vm)?;
replaced.drop_with(vm);
Ok(())
}
const UNSUPPORTED_MEMBERS: [(&str, &str); 1] = [("__post_init__", "which would be silently skipped")];
fn reject_unsupported_members<'h>(class: &HeapRead<'h, Class>, vm: &VM<'h>) -> RunResult<()> {
let namespace = class.get(vm.heap).namespace();
match UNSUPPORTED_MEMBERS
.iter()
.find(|&&(name, _)| namespace.get_by_str(name, vm.heap, vm.interns).is_some())
{
Some((name, consequence)) => Err(ExcType::not_implemented(format!(
"dataclass() does not yet support {name} in a class body, {consequence}"
))
.into()),
None => Ok(()),
}
}
fn collect_annotated_fields<'h>(class: &HeapRead<'h, Class>, vm: &VM<'h>) -> Vec<DataclassField> {
let namespace = class.get(vm.heap).namespace();
let ann_id = match namespace.get_by_str("__annotations__", vm.heap, vm.interns) {
Some(Value::Ref(id)) => *id,
_ => return Vec::new(),
};
let HeapData::Dict(annotations) = vm.heap.get(ann_id) else {
return Vec::new();
};
let mut fields = Vec::new();
for (key, annotation) in annotations {
let Value::InternString(name_id) = key else { continue };
if is_classvar(&annotation_text(annotation, vm)) {
continue;
}
let default = namespace
.get_by_str(vm.interns.get_str(*name_id), vm.heap, vm.interns)
.map(|v| v.clone_with_heap(vm.heap));
fields.push(DataclassField::new(
*name_id,
annotation.clone_with_heap(vm.heap),
default,
));
}
fields
}
fn annotation_text(annotation: &Value, vm: &VM<'_>) -> String {
annotation
.as_either_str(vm.heap)
.map(|s| s.as_str(vm.interns).to_owned())
.unwrap_or_default()
}
fn validate_fields(vm: &mut VM<'_>, fields: &[DataclassField]) -> RunResult<()> {
for field in fields {
let name = vm.interns.get_str(field.name()).to_owned();
if is_initvar(&annotation_text(field.annotation(), vm)) {
return Err(ExcType::not_implemented(format!(
"dataclass() does not yet support InitVar (field {name}), which would become an ordinary field"
))
.into());
}
if let Some(default) = field.default() {
if default.py_hash(vm)?.is_none() {
let ty = default.py_type_name(vm);
return Err(ExcType::value_error(format!(
"mutable default <class '{ty}'> for field {name} is not allowed: use default_factory"
)));
}
}
}
if let Some((prev, field)) = first_non_default_after_default(fields) {
let (prev, field) = (vm.interns.get_str(prev), vm.interns.get_str(field));
return Err(ExcType::type_error(format!(
"non-default argument '{field}' follows default argument '{prev}'"
)));
}
Ok(())
}
fn first_non_default_after_default(fields: &[DataclassField]) -> Option<(StringId, StringId)> {
let mut last_default = None;
fields
.iter()
.find_map(|field| match (field.default().is_some(), last_default) {
(true, _) => {
last_default = Some(field.name());
None
}
(false, Some(prev)) => Some((prev, field.name())),
(false, None) => None,
})
}
fn is_classvar(annotation: &str) -> bool {
annotation_head(annotation, "ClassVar")
}
fn is_initvar(annotation: &str) -> bool {
annotation_head(annotation, "InitVar")
}
fn annotation_head(annotation: &str, name: &str) -> bool {
let s = annotation.trim().trim_matches(|c| c == '\'' || c == '"').trim();
let head = s.split_once('[').map_or(s, |(head, _)| head).trim_end();
head.rsplit('.').next() == Some(name)
}
fn fields_dict_id(namespace: &Dict, vm: &VM<'_>) -> Option<HeapId> {
match namespace.get_by_str(StaticStrings::DataclassFields.into(), vm.heap, vm.interns) {
Some(Value::Ref(id)) if matches!(vm.heap.get(*id), HeapData::Dict(_)) => Some(*id),
_ => None,
}
}
fn class_fields_dict_id(class_id: HeapId, vm: &VM<'_>) -> Option<HeapId> {
match vm.heap.get(class_id) {
HeapData::Class(class) => fields_dict_id(class.namespace(), vm),
_ => None,
}
}
pub(crate) fn dataclass_fields(class_id: HeapId, vm: &VM<'_>) -> Option<Vec<StringId>> {
let fields_id = class_fields_dict_id(class_id, vm)?;
Some(field_specs(fields_id, vm).into_iter().map(|(name, _)| name).collect())
}
fn field_specs(fields_id: HeapId, vm: &VM<'_>) -> Vec<(StringId, bool)> {
let HeapData::Dict(fields) = vm.heap.get(fields_id) else {
return Vec::new();
};
fields
.iter()
.filter_map(|(_, value)| match value {
Value::Ref(id) => match vm.heap.get(*id) {
HeapData::DataclassField(field) => Some((field.name(), field.default().is_some())),
_ => None,
},
_ => None,
})
.collect()
}
pub(crate) fn is_dataclass_class(class_id: HeapId, vm: &VM<'_>) -> bool {
class_fields_dict_id(class_id, vm).is_some()
}
pub(crate) fn dataclass_init<'h>(
vm: &mut VM<'h>,
class: &HeapRead<'h, Class>,
instance: Value,
args: ArgValues,
) -> Result<CallResult, RunError> {
let fields_id =
fields_dict_id(class.get(vm.heap).namespace(), vm).expect("dataclass_init requires __dataclass_fields__");
let fields = field_specs(fields_id, vm);
let mut guard = DropGuard::new(instance, vm);
let (instance, vm) = guard.as_parts();
let Value::Ref(instance_id) = instance else {
unreachable!("the caller just allocated the instance")
};
let HeapReadOutput::Instance(mut instance) = vm.heap.read(*instance_id) else {
unreachable!("the caller just allocated the instance")
};
let values = bind_dataclass_fields(vm, class, fields_id, &fields, args)?;
store_bound_fields(&mut instance, &fields, values, vm)?;
Ok(CallResult::Value(guard.into_inner()))
}
fn store_bound_fields<'h>(
instance: &mut HeapRead<'h, Instance>,
fields: &[(StringId, bool)],
values: Vec<Value>,
vm: &mut VM<'h>,
) -> Result<(), RunError> {
defer_drop_mut!(values, vm);
for i in 0..values.len() {
let value = mem::replace(&mut values[i], Value::None);
let name = Value::InternString(fields[i].0);
let replaced = instance.set_attr(name, value, vm)?;
replaced.drop_with(vm);
}
Ok(())
}
fn bind_dataclass_fields<'h>(
vm: &mut VM<'h>,
class: &HeapRead<'h, Class>,
fields_id: HeapId,
fields: &[(StringId, bool)],
args: ArgValues,
) -> Result<Vec<Value>, RunError> {
let n_fields = fields.len();
let init_name = format!("{}.__init__", class.get(vm.heap).name().as_str(vm.interns));
let (pos_iter, kwargs) = args.into_parts();
let values: Vec<Option<Value>> = (0..n_fields).map(|_| None).collect();
let mut guard = DropGuard::new(values, vm);
let (values, vm) = guard.as_parts_mut();
let mut n_pos = 0;
for (slot, value) in pos_iter.enumerate() {
n_pos += 1;
match values.get_mut(slot) {
Some(unbound) => *unbound = Some(value),
None => value.drop_with(vm),
}
}
bind_keyword_args(values, fields, kwargs, &init_name, vm)?;
if n_pos > n_fields {
let required = fields.iter().filter(|(_, has_default)| !has_default).count();
return Err(ExcType::type_error_too_many_positional_range(
&init_name,
1 + required,
1 + n_fields,
1 + n_pos,
0,
));
}
let mut missing: Vec<String> = Vec::new();
for (idx, (id, _)) in fields.iter().enumerate() {
if values[idx].is_some() {
continue;
}
match captured_default(vm, fields_id, idx) {
Some(default) => values[idx] = Some(default),
None => missing.push(vm.interns.get_str(*id).to_owned()),
}
}
if !missing.is_empty() {
let refs: Vec<&str> = missing.iter().map(String::as_str).collect();
return Err(ExcType::type_error_missing_positional_with_names(&init_name, &refs));
}
Ok(guard
.into_inner()
.into_iter()
.map(|v| v.expect("all fields bound"))
.collect())
}
fn bind_keyword_args(
values: &mut [Option<Value>],
fields: &[(StringId, bool)],
kwargs: KwargsValues,
init_name: &str,
vm: &mut VM<'_>,
) -> RunResult<()> {
let kwargs = kwargs.into_iter();
defer_drop_mut!(kwargs, vm);
for (key, value) in kwargs.by_ref() {
defer_drop!(key, vm);
let name = key
.as_either_str(vm.heap)
.expect("DictMerge rejects non-string keys before the call")
.as_str(vm.interns)
.to_owned();
let slot = fields.iter().position(|(id, _)| vm.interns.get_str(*id) == name);
let value = DropGuard::new(value, vm);
match slot {
Some(idx) if values[idx].is_none() => values[idx] = Some(value.into_inner()),
Some(_) => return Err(ExcType::type_error_duplicate_arg(init_name, &name)),
None => return Err(ExcType::type_error_unexpected_keyword(init_name, &name)),
}
}
Ok(())
}
fn captured_default(vm: &VM<'_>, fields_id: HeapId, idx: usize) -> Option<Value> {
let HeapData::Dict(fields) = vm.heap.get(fields_id) else {
return None;
};
match fields.value_at(idx) {
Some(Value::Ref(id)) => match vm.heap.get(*id) {
HeapData::DataclassField(field) => field.default().map(|v| v.clone_with_heap(vm.heap)),
_ => None,
},
_ => None,
}
}
pub(crate) fn dataclass_eq(
self_id: HeapId,
field_names: &[StringId],
other: &Value,
vm: &mut VM<'_>,
) -> RunResult<Option<bool>> {
let class_id = instance_class(self_id, vm);
let &Value::Ref(other_id) = other else {
return Ok(None);
};
if !matches!(vm.heap.get(other_id), HeapData::Instance(inst) if inst.class() == class_id) {
return Ok(None);
}
let mut guard = vm.recursion_guard()?;
let vm = &mut *guard;
for name_id in field_names {
let field_name = vm.interns.get_str(*name_id).to_owned();
let a = instance_attr(self_id, &field_name, vm);
defer_drop!(a, vm);
let b = instance_attr(other_id, &field_name, vm);
defer_drop!(b, vm);
match (a, b) {
(Some(a), Some(b)) if !a.py_eq_operator(b, vm)? => return Ok(Some(false)),
(Some(_), Some(_)) => {}
_ => {
let class_name = class_name(class_id, vm.heap, vm.interns).into_owned();
return Err(ExcType::attribute_error(&class_name, &field_name));
}
}
}
Ok(Some(true))
}
pub(crate) fn dataclass_repr_fmt(
self_id: HeapId,
field_names: &[StringId],
f: &mut impl Write,
vm: &mut VM<'_>,
heap_ids: &mut LazyHeapSet,
) -> RunResult<()> {
let class_id = instance_class(self_id, vm);
let name = class_name(class_id, vm.heap, vm.interns).to_string();
write_dataclass_repr(f, &name, field_names.len(), vm, heap_ids, |i, vm| {
let field_name = vm.interns.get_str(field_names[i]).to_owned();
match instance_attr(self_id, &field_name, vm) {
Some(value) => Ok((field_name, Some(value))),
None => Err(ExcType::attribute_error(&name, &field_name)),
}
})
}
fn instance_class(self_id: HeapId, vm: &VM<'_>) -> HeapId {
match vm.heap.get(self_id) {
HeapData::Instance(inst) => inst.class(),
_ => unreachable!("dataclass dispatch on a non-instance"),
}
}
fn is_dataclass(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let arg = args.get_one_arg("is_dataclass", vm.heap)?;
let result = match &arg {
Value::Ref(id) => match vm.heap.get(*id) {
HeapData::Class(_) => is_dataclass_class(*id, vm),
HeapData::Instance(instance) => is_dataclass_class(instance.class(), vm),
_ => false,
},
_ => false,
};
arg.drop_with(vm);
Ok(Value::Bool(result))
}