use std::{
fmt::Write,
hash::{DefaultHasher, Hash, Hasher},
mem,
};
use serde::ser::SerializeStruct;
use super::{Dict, LazyHeapSet, PyTrait};
use crate::{
args::ArgValues,
bytecode::{CallResult, VM},
defer_drop,
exception_private::{ExcType, RunResult, SimpleException},
hash::HashValue,
heap::{
BorrowedHeapRead, BorrowedHeapReadMut, HeapId, HeapItem, HeapRead, HeapReadOutput, heap_read_ref_as_field,
heap_read_ref_as_field_mut,
},
intern::Interns,
resource::ResourceTracker,
types::Type,
value::{EitherStr, Value},
};
#[derive(Debug)]
pub(crate) struct Dataclass {
name: EitherStr,
type_id: u64,
field_names: Vec<String>,
attrs: Dict,
frozen: bool,
}
impl Dataclass {
#[must_use]
pub fn new(name: impl Into<EitherStr>, type_id: u64, field_names: Vec<String>, attrs: Dict, frozen: bool) -> Self {
Self {
name: name.into(),
type_id,
field_names,
attrs,
frozen,
}
}
#[must_use]
pub fn name<'a>(&'a self, interns: &'a Interns) -> &'a str {
self.name.as_str(interns)
}
#[must_use]
pub fn type_id(&self) -> u64 {
self.type_id
}
#[must_use]
pub fn field_names(&self) -> &[String] {
&self.field_names
}
#[must_use]
pub fn attrs(&self) -> &Dict {
&self.attrs
}
#[must_use]
pub fn is_frozen(&self) -> bool {
self.frozen
}
}
impl<'h> HeapRead<'h, Dataclass> {
pub fn set_attr(
&mut self,
name: Value,
value: Value,
vm: &mut VM<'h, impl ResourceTracker>,
) -> RunResult<Option<Value>> {
if self.get(vm.heap).frozen {
let name_repr = name.py_repr(vm)?;
defer_drop!(name_repr, vm);
let exc = SimpleException::new_msg(
ExcType::FrozenInstanceError,
format!("cannot assign to field {}", name_repr.to_str(vm)?),
);
name.drop_with_heap(vm);
value.drop_with_heap(vm);
return Err(exc.into());
}
self.attrs_mut().set(name, value, vm)
}
pub fn attrs(&self) -> BorrowedHeapRead<'_, 'h, Dict> {
heap_read_ref_as_field!(self, Dataclass, attrs)
}
pub fn attrs_mut(&mut self) -> BorrowedHeapReadMut<'_, 'h, Dict> {
heap_read_ref_as_field_mut!(self, Dataclass, attrs)
}
}
impl<'h> PyTrait<'h> for HeapRead<'h, Dataclass> {
fn py_type(&self, _vm: &VM<'h, impl ResourceTracker>) -> Type {
Type::Dataclass
}
fn py_len(&self, _vm: &VM<'h, impl ResourceTracker>) -> Option<usize> {
None
}
fn py_eq_impl(&self, other: &Value, vm: &mut VM<'h, impl ResourceTracker>) -> RunResult<Option<bool>> {
let Some(HeapReadOutput::Dataclass(other)) = other.read_heap(vm) else {
return Ok(None);
};
if self.get(vm.heap).type_id() != other.get(vm.heap).type_id() {
return Ok(Some(false));
}
Ok(Some(self.attrs().eq_dict(&other.attrs(), vm)?))
}
fn py_hash(&self, _self_id: HeapId, vm: &mut VM<'h, impl ResourceTracker>) -> RunResult<Option<HashValue>> {
if !self.get(vm.heap).frozen {
return Ok(None);
}
let mut guard = vm.recursion_guard()?;
let vm = &mut *guard;
let mut hasher = DefaultHasher::new();
self.get(vm.heap).name.hash(&mut hasher);
let field_count = self.get(vm.heap).field_names.len();
for i in 0..field_count {
let field_name = &self.get(vm.heap).field_names[i];
field_name.hash(&mut hasher);
if let Some(value) = self.get(vm.heap).attrs.get_by_str(field_name, vm.heap, vm.interns) {
let value = value.clone_with_heap(vm.heap);
defer_drop!(value, vm);
match value.py_hash(vm)? {
Some(h) => h.hash(&mut hasher),
None => return Ok(None),
}
}
}
Ok(Some(HashValue::new(hasher.finish())))
}
fn py_bool(&self, _vm: &mut VM<'h, impl ResourceTracker>) -> bool {
true
}
fn py_repr_fmt(
&self,
f: &mut impl Write,
vm: &mut VM<'h, impl ResourceTracker>,
heap_ids: &mut LazyHeapSet,
) -> RunResult<()> {
let Ok(mut guard) = vm.recursion_guard() else {
return Ok(f.write_str("...")?);
};
let vm = &mut *guard;
let dc = self.get(vm.heap);
f.write_str(dc.name(vm.interns))?;
f.write_char('(')?;
let field_count = self.get(vm.heap).field_names.len();
let interns = vm.interns;
for i in 0..field_count {
if i > 0 {
f.write_str(", ")?;
}
let field_name = &self.get(vm.heap).field_names[i];
f.write_str(field_name)?;
f.write_char('=')?;
if let Some(value) = self.get(vm.heap).attrs.get_by_str(field_name, vm.heap, interns) {
let value = value.clone_with_heap(vm.heap);
defer_drop!(value, vm);
value.py_repr_fmt(f, vm, heap_ids)?;
} else {
f.write_str("<?>")?;
}
}
f.write_char(')')?;
Ok(())
}
fn py_call_attr(
&mut self,
self_id: HeapId,
vm: &mut VM<'h, impl ResourceTracker>,
attr: &EitherStr,
args: ArgValues,
) -> RunResult<CallResult> {
let attr_str = attr.as_str(vm.interns);
if !attr_str.starts_with('_')
&& self
.get(vm.heap)
.attrs
.get_by_str(attr_str, vm.heap, vm.interns)
.is_none()
{
vm.heap.inc_ref(self_id);
let self_arg = Value::Ref(self_id);
let args_with_self = args.prepend(self_arg);
Ok(CallResult::MethodCall(attr.clone(), args_with_self))
} else {
let method_name = attr.as_str(vm.interns);
defer_drop!(args, vm);
if let Some(value) = self.get(vm.heap).attrs.get_by_str(method_name, vm.heap, vm.interns) {
let type_name = value.py_type_name(vm);
Err(ExcType::type_error_not_callable_object(&type_name))
} else {
Err(ExcType::attribute_error(
self.get(vm.heap).name(vm.interns),
method_name,
))
}
}
}
fn py_getattr(&self, attr: &EitherStr, vm: &mut VM<'h, impl ResourceTracker>) -> RunResult<Option<CallResult>> {
let attr_name = attr.as_str(vm.interns);
match self.get(vm.heap).attrs.get_by_str(attr_name, vm.heap, vm.interns) {
Some(value) => Ok(Some(CallResult::Value(value.clone_with_heap(vm.heap)))),
None => Err(ExcType::attribute_error(self.get(vm.heap).name(vm.interns), attr_name)),
}
}
}
impl HeapItem for Dataclass {
fn py_estimate_size(&self) -> usize {
mem::size_of::<Self>()
+ self.name.py_estimate_size()
+ self.field_names.iter().map(String::len).sum::<usize>()
+ self.attrs.py_estimate_size()
}
fn py_dec_ref_ids(&mut self, stack: &mut Vec<HeapId>) {
self.attrs.py_dec_ref_ids(stack);
}
}
impl serde::Serialize for Dataclass {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut state = serializer.serialize_struct("Dataclass", 5)?;
state.serialize_field("name", &self.name)?;
state.serialize_field("type_id", &self.type_id)?;
state.serialize_field("field_names", &self.field_names)?;
state.serialize_field("attrs", &self.attrs)?;
state.serialize_field("frozen", &self.frozen)?;
state.end()
}
}
impl<'de> serde::Deserialize<'de> for Dataclass {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
#[derive(serde::Deserialize)]
struct DataclassData {
name: EitherStr,
type_id: u64,
field_names: Vec<String>,
attrs: Dict,
frozen: bool,
}
let dc = DataclassData::deserialize(deserializer)?;
Ok(Self {
name: dc.name,
type_id: dc.type_id,
field_names: dc.field_names,
attrs: dc.attrs,
frozen: dc.frozen,
})
}
}