diff --git a/cpp/csp/python/PyStructList.hi b/cpp/csp/python/PyStructList.hi index bf829ea6c..7e09a5a2d 100644 --- a/cpp/csp/python/PyStructList.hi +++ b/cpp/csp/python/PyStructList.hi @@ -394,6 +394,16 @@ static PyMappingMethods py_struct_list_as_mapping = { py_struct_list_ass_subscript }; +static PyObject * +PyStructList_new( PyTypeObject *type, PyObject *args, PyObject *kwds ) +{ + // Since the PyStructList has no real meaning when created from Python, we can reconstruct the PSL's value + // by just treating it as a list. Thus, we simply override the tp_new behaviour to return a list object here. + // Again, since we don't have tp_init for the PSL, we need to rely on the Python list's tp_init function. + + return PyObject_Call( ( PyObject * ) &PyList_Type, args, kwds ); // Calls both tp_new and tp_init for a Python list +} + template static int PyStructList_tp_clear( PyStructList * self ) @@ -437,7 +447,7 @@ template PyTypeObject PyStructList::PyType = { .tp_clear = ( inquiry ) PyStructList_tp_clear, .tp_methods = PyStructList_methods, .tp_alloc = PyType_GenericAlloc, - .tp_new = PyType_GenericNew, + .tp_new = PyStructList_new, .tp_free = PyObject_GC_Del, }; diff --git a/csp/impl/struct.py b/csp/impl/struct.py index bf2206aa5..89eaa254b 100644 --- a/csp/impl/struct.py +++ b/csp/impl/struct.py @@ -108,9 +108,7 @@ def _obj_to_python(cls, obj): ) elif isinstance(obj, dict): return {k: cls._obj_to_python(v) for k, v in obj.items()} - elif isinstance(obj, list): - return list(cls._obj_to_python(v) for v in obj) - elif isinstance(obj, (tuple, set)): + elif isinstance(obj, (list, tuple, set)): return type(obj)(cls._obj_to_python(v) for v in obj) elif isinstance(obj, csp.Enum): return obj.name # handled in _obj_from_python diff --git a/csp/tests/impl/test_struct.py b/csp/tests/impl/test_struct.py index b5e634a9d..0920aa5cc 100644 --- a/csp/tests/impl/test_struct.py +++ b/csp/tests/impl/test_struct.py @@ -697,6 +697,19 @@ def test_from_dict_with_enum(self): struct = StructWithDefaults.from_dict({"e": MyEnum.A}) self.assertEqual(MyEnum.A, getattr(struct, "e")) + def test_from_dict_with_list_derived_type(self): + class ListDerivedType(list): + def __init__(self, iterable=None): + super().__init__(iterable) + + class StructWithListDerivedType(csp.Struct): + ldt: ListDerivedType + + s1 = StructWithListDerivedType(ldt=ListDerivedType([1,2])) + self.assertTrue(isinstance(s1.to_dict()['ldt'], ListDerivedType)) + s2 = StructWithListDerivedType.from_dict(s1.to_dict()) + self.assertEqual(s1, s2) + def test_from_dict_loop_no_defaults(self): looped = StructNoDefaults.from_dict(StructNoDefaults(a1=[9, 10]).to_dict()) self.assertEqual(looped, StructNoDefaults(a1=[9, 10]))