use std::{
collections::hash_map::DefaultHasher,
hash::{Hash, Hasher},
};
use monty_types::{DictPairs, MontyObject};
use pyo3::{
Bound,
exceptions::{PyAttributeError, PyTypeError},
intern,
prelude::*,
sync::PyOnceLock,
types::{PyDict, PyList, PyString, PyType},
};
use super::convert::{monty_to_py_inner, py_to_monty};
pub fn is_dataclass(value: &Bound<'_, PyAny>) -> bool {
value
.hasattr(intern!(value.py(), "__dataclass_fields__"))
.unwrap_or(false)
&& !value.is_instance_of::<PyType>()
}
pub fn dataclass_to_monty(value: &Bound<'_, PyAny>, dc_registry: &DcRegistry, depth: u8) -> PyResult<MontyObject> {
let py = value.py();
let dc_type = value.get_type();
let name: String = dc_type.getattr(intern!(py, "__name__"))?.extract()?;
let type_id = dc_type.as_ptr() as u64;
let fields_dict = value
.getattr(intern!(py, "__dataclass_fields__"))?
.cast_into::<PyDict>()?;
let frozen = value
.getattr(intern!(py, "__dataclass_params__"))?
.getattr(intern!(py, "frozen"))?
.extract::<bool>()?;
let field_type_marker = get_field_marker(py)?;
let mut field_names = Vec::new();
let mut attrs = Vec::new();
for (field_name_obj, field) in fields_dict.iter() {
let field_type = field.getattr(intern!(py, "_field_type"))?;
if field_type.is(field_type_marker) {
let field_name_str = field_name_obj.cast::<PyString>()?.to_str()?.to_string();
if field_name_str.starts_with('_') {
continue;
}
let field_value = value.getattr(field_name_obj.cast::<PyString>()?)?;
let field_name_monty = py_to_monty(&field_name_obj, dc_registry, depth)?;
let field_value_monty = py_to_monty(&field_value, dc_registry, depth)?;
field_names.push(field_name_str);
attrs.push((field_name_monty, field_value_monty));
}
}
Ok(MontyObject::Dataclass {
name,
type_id,
field_names,
attrs: attrs.into(),
frozen,
})
}
#[expect(clippy::too_many_arguments)]
pub fn dataclass_to_py(
py: Python<'_>,
name: &str,
type_id: u64,
field_names: &[String],
attrs: &DictPairs,
frozen: bool,
dc_registry: &DcRegistry,
depth: u8,
) -> PyResult<Py<PyAny>> {
if let Some(original_type_py) = dc_registry.get(py, type_id)? {
let original_type = original_type_py.bind(py).cast::<PyType>()?;
let kwargs = PyDict::new(py);
for (key, value) in attrs {
if let MontyObject::String(s) = key {
let key_str = s.as_str();
if field_names.iter().any(|f| f.as_str() == key_str) {
kwargs.set_item(key_str, monty_to_py_inner(py, value, dc_registry, depth)?)?;
}
}
}
original_type.call((), Some(&kwargs)).map(Bound::unbind)
} else {
let dc = PyUnknownDataclass::new(
py,
name.to_string(),
field_names.to_vec(),
attrs,
frozen,
dc_registry,
depth,
)?;
Ok(Py::new(py, dc)?.into_any())
}
}
#[derive(Debug)]
pub struct DcRegistry {
registry: Py<PyDict>,
}
impl DcRegistry {
#[must_use]
pub fn new(py: Python<'_>) -> Self {
Self {
registry: PyDict::new(py).unbind(),
}
}
pub fn from_list(py: Python<'_>, dataclass_registry: Option<&Bound<'_, PyList>>) -> PyResult<Self> {
let slf = Self::new(py);
if let Some(registry_list) = dataclass_registry {
for cls in registry_list {
slf.insert(&cls)?;
}
}
Ok(slf)
}
#[must_use]
pub fn clone_ref(&self, py: Python<'_>) -> Self {
Self {
registry: self.registry.clone_ref(py),
}
}
pub fn insert<T>(&self, obj: &Bound<'_, T>) -> PyResult<()> {
let py = obj.py();
let type_id = obj.as_ptr() as u64;
self.registry.bind(py).set_item(type_id, obj.as_any())
}
pub fn get(&self, py: Python<'_>, type_id: u64) -> PyResult<Option<Py<PyAny>>> {
Ok(self.registry.bind(py).get_item(type_id)?.map(Bound::unbind))
}
}
#[pyclass(name = "UnknownDataclass")]
pub struct PyUnknownDataclass {
name: String,
field_names: Vec<String>,
attrs: Py<PyDict>,
frozen: bool,
}
#[pymethods]
impl PyUnknownDataclass {
#[getter]
fn __dataclass_fields__(&self, py: Python<'_>) -> PyResult<Py<PyDict>> {
let field_marker = get_field_marker(py)?;
let missing = get_missing(py)?;
let field_class = get_field_class(py)?;
let attrs = self.attrs.bind(py);
let fields_dict = PyDict::new(py);
for field_name in &self.field_names {
let field_type = if let Some(value) = attrs.get_item(field_name)? {
value.get_type().into_any()
} else {
py.None().into_bound(py).get_type().into_any()
};
let field_obj = if py.version_info() >= (3, 14) {
field_class.call1((
missing, missing, true, true, py.None(), true, py.None(), false, py.None(), ))?
} else {
field_class.call1((
missing, missing, true, true, py.None(), true, py.None(), false, ))?
};
field_obj.setattr("name", field_name)?;
field_obj.setattr("type", field_type)?;
field_obj.setattr("_field_type", field_marker)?;
fields_dict.set_item(field_name, field_obj)?;
}
Ok(fields_dict.unbind())
}
#[getter]
fn __dataclass_params__(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
let params_class = get_dataclass_params_class(py)?;
let params = if py.version_info() >= (3, 12) {
params_class.call1((
true, true, true, false, false, self.frozen, true, false, false, false, ))?
} else {
params_class.call1((
true, true, true, false, false, self.frozen, ))?
};
Ok(params.unbind())
}
fn __getattr__(&self, py: Python<'_>, name: &str) -> PyResult<Py<PyAny>> {
let attrs = self.attrs.bind(py);
match attrs.get_item(name)? {
Some(value) => Ok(value.unbind()),
None => Err(PyAttributeError::new_err(format!(
"'UnknownDataclass' object has no attribute '{name}'",
))),
}
}
fn __setattr__(&self, py: Python<'_>, name: &str, value: Py<PyAny>) -> PyResult<()> {
if self.frozen {
let frozen_error = get_frozen_instance_error(py)?;
let msg = format!("cannot assign to field '{name}'");
return Err(PyErr::from_value(frozen_error.call1((msg,))?));
}
let attrs = self.attrs.bind(py);
attrs.set_item(name, value)?;
Ok(())
}
fn __repr__(&self, py: Python<'_>) -> PyResult<String> {
let attrs = self.attrs.bind(py);
let mut parts = Vec::new();
for field_name in &self.field_names {
if let Some(value) = attrs.get_item(field_name)? {
let value_repr: String = value.repr()?.extract()?;
parts.push(format!("{field_name}={value_repr}"));
}
}
Ok(format!("<Unknown Dataclass {}({})>", self.name, parts.join(", ")))
}
fn __eq__(&self, py: Python<'_>, other: &Bound<'_, PyAny>) -> PyResult<bool> {
if let Ok(other_dc) = other.extract::<PyRef<'_, Self>>() {
if self.name != other_dc.name {
return Ok(false);
}
let self_attrs = self.attrs.bind(py);
let other_attrs = other_dc.attrs.bind(py);
self_attrs.eq(other_attrs)
} else {
Ok(false)
}
}
fn __hash__(&self, py: Python<'_>) -> PyResult<isize> {
if !self.frozen {
return Err(PyTypeError::new_err("unhashable type: 'UnknownDataclass'"));
}
let mut hasher = DefaultHasher::new();
let attrs = self.attrs.bind(py);
for field_name in &self.field_names {
field_name.hash(&mut hasher);
if let Some(value) = attrs.get_item(field_name)? {
let value_hash: isize = value.hash()?;
value_hash.hash(&mut hasher);
}
}
let hash_u64 = hasher.finish();
#[cfg(target_pointer_width = "64")]
let hash_isize = isize::from_ne_bytes(hash_u64.to_ne_bytes());
#[cfg(not(target_pointer_width = "64"))]
let hash_isize = {
let hash_u32 = hash_u64 as u32;
i32::from_ne_bytes(hash_u32.to_ne_bytes()) as isize
};
Ok(hash_isize)
}
}
impl PyUnknownDataclass {
pub fn new<'a>(
py: Python<'_>,
name: String,
field_names: Vec<String>,
attrs: impl IntoIterator<Item = &'a (MontyObject, MontyObject)>,
frozen: bool,
dc_registry: &DcRegistry,
depth: u8,
) -> PyResult<Self> {
let dict = PyDict::new(py);
for (k, v) in attrs {
dict.set_item(
monty_to_py_inner(py, k, dc_registry, depth)?,
monty_to_py_inner(py, v, dc_registry, depth)?,
)?;
}
Ok(Self {
name,
field_names,
attrs: dict.unbind(),
frozen,
})
}
}
fn get_field_marker(py: Python<'_>) -> PyResult<&Bound<'_, PyAny>> {
static DC_FIELD_MARKER: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
DC_FIELD_MARKER.import(py, "dataclasses", "_FIELD")
}
fn get_missing(py: Python<'_>) -> PyResult<&Bound<'_, PyAny>> {
static DC_MISSING: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
DC_MISSING.import(py, "dataclasses", "MISSING")
}
fn get_field_class(py: Python<'_>) -> PyResult<&Bound<'_, PyAny>> {
static DC_FIELD_CLASS: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
DC_FIELD_CLASS.import(py, "dataclasses", "Field")
}
fn get_dataclass_params_class(py: Python<'_>) -> PyResult<&Bound<'_, PyAny>> {
static DC_PARAMS_CLASS: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
DC_PARAMS_CLASS.import(py, "dataclasses", "_DataclassParams")
}
pub fn get_frozen_instance_error(py: Python<'_>) -> PyResult<&Bound<'_, PyAny>> {
static DC_FROZEN_ERROR: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
DC_FROZEN_ERROR.import(py, "dataclasses", "FrozenInstanceError")
}