2323if TYPE_CHECKING :
2424 from collections .abc import Callable , Hashable , Iterable
2525
26+ from _typeshed import DataclassInstance
2627 from sympy .printing .latex import LatexPrinter
2728
2829 if sys .version_info >= (3 , 11 ):
@@ -279,7 +280,7 @@ def new_method(cls, *args, evaluate: bool = False, **kwargs) -> type[ExprClass]:
279280 return expr
280281
281282 cls .__new__ = new_method
282- cls .__getnewargs__ = dataclasses . astuple
283+ cls .__getnewargs__ = _get_field_values
283284 cls ._hashable_content = _hashable_content_method
284285 if non_sympy_fields :
285286 cls ._eval_subs = _eval_subs_method
@@ -486,7 +487,7 @@ def class_wrapper(cls: T) -> T:
486487def _eval_subs_method (self , old , new , ** hints ):
487488 # https://github.com/sympy/sympy/blob/1.12/sympy/core/basic.py#L1117-L1147
488489 hit = False
489- old_args = dataclasses . astuple (self )
490+ old_args = _get_field_values (self )
490491 new_args = list (old_args )
491492 for i , old_arg in enumerate (old_args ):
492493 if not hasattr (old_arg , "_eval_subs" ):
@@ -535,7 +536,7 @@ def _xreplace_method(self, rule) -> tuple[sp.Expr, bool]:
535536 if rule :
536537 new_args = []
537538 hit = False
538- for arg in dataclasses . astuple (self ):
539+ for arg in _get_field_values (self ):
539540 if hasattr (arg , "_xreplace" ) and not isclass (arg ):
540541 replace_result , is_replaced = arg ._xreplace (rule ) # ruff: ignore[private-member-access]
541542 elif isinstance (rule , abc .Mapping ):
@@ -551,13 +552,28 @@ def _xreplace_method(self, rule) -> tuple[sp.Expr, bool]:
551552 return self , False
552553
553554
554- def get_sympy_fields (cls ) -> tuple [Field , ...]:
555+ def _get_field_values (self : DataclassInstance ) -> tuple [Any , ...]:
556+ """Get the field values of a dataclass-like class without recursing.
557+
558+ This is a shallow alternative to :func:`dataclasses.astuple`, which recurses into
559+ nested dataclasses and converts them to `tuple` as well. Since decorated classes are
560+ both dataclasses and `~sympy.core.expr.Expr` instances, that recursion would destroy
561+ nested expressions, for instance when pickling.
562+ """
563+ return tuple (getattr (self , field .name ) for field in dataclasses .fields (self ))
564+
565+
566+ def get_sympy_fields (
567+ cls : DataclassInstance | type [DataclassInstance ],
568+ ) -> tuple [Field [Any ], ...]:
555569 return tuple (f for f in dataclasses .fields (cls ) if _is_sympify (f ))
556570
557571
558- def get_non_sympy_fields (cls ) -> tuple [Field , ...]:
572+ def get_non_sympy_fields (
573+ cls : DataclassInstance | type [DataclassInstance ],
574+ ) -> tuple [Field [Any ], ...]:
559575 return tuple (f for f in dataclasses .fields (cls ) if not _is_sympify (f ))
560576
561577
562- def _is_sympify (field : Field ) -> bool :
578+ def _is_sympify (field : Field [ Any ] ) -> bool :
563579 return bool (field .metadata .get ("sympify" ))
0 commit comments