fix(numpy): accept PYBIND11_TYPE-wrapped types in the dtype macros (#6148)

* fix(numpy): accept PYBIND11_TYPE-wrapped types in the dtype macros

* docs: point at the macro notes from the structured types section

* fix(numpy): accept bare types in direct field descriptor macro calls
This commit is contained in:
Advit Arora
2026-08-19 16:05:59 -07:00
committed by GitHub
parent a49c4dbfa4
commit 22bf5b6ed4
4 changed files with 62 additions and 7 deletions
+4
View File
@@ -232,6 +232,10 @@ prevent many types of unsupported structures, it is still the user's
responsibility to use only "plain" structures that can be safely manipulated as
raw memory without violating invariants.
Types whose spelling contains a comma must be wrapped in ``PYBIND11_TYPE``:
``PYBIND11_NUMPY_DTYPE(PYBIND11_TYPE(C<int, double>), x, y)``.
See :ref:`macro_notes`.
Scalar types
============
+20 -7
View File
@@ -1799,16 +1799,29 @@ private:
# define PYBIND11_NUMPY_DTYPE_EX(Type, ...) ((void) 0)
#else
# define PYBIND11_FIELD_DESCRIPTOR_EX(T, Field, Name) \
// The _IMPL variants take T parenthesized to survive the comma re-splitting in the
// PYBIND11_MAP_LIST expansions below; the plain variants keep accepting a bare type (see #4018).
# define PYBIND11_UNPAREN_TYPE(T) PYBIND11_TYPE T
# define PYBIND11_FIELD_DESCRIPTOR_EX_IMPL(T, Field, Name) \
::pybind11::detail::field_descriptor { \
Name, offsetof(T, Field), sizeof(decltype(std::declval<T>().Field)), \
::pybind11::format_descriptor<decltype(std::declval<T>().Field)>::format(), \
Name, offsetof(PYBIND11_UNPAREN_TYPE(T), Field), \
sizeof(decltype(std::declval<PYBIND11_UNPAREN_TYPE(T)>().Field)), \
::pybind11::format_descriptor< \
decltype(std::declval<PYBIND11_UNPAREN_TYPE(T)>().Field)>::format(), \
::pybind11::detail::npy_format_descriptor< \
decltype(std::declval<T>().Field)>::dtype() \
decltype(std::declval<PYBIND11_UNPAREN_TYPE(T)>().Field)>::dtype() \
}
# define PYBIND11_FIELD_DESCRIPTOR_IMPL(T, Field) \
PYBIND11_FIELD_DESCRIPTOR_EX_IMPL(T, Field, #Field)
# define PYBIND11_FIELD_DESCRIPTOR_EX(T, Field, Name) \
PYBIND11_FIELD_DESCRIPTOR_EX_IMPL((T), Field, Name)
// Extract name, offset and format descriptor for a struct field
# define PYBIND11_FIELD_DESCRIPTOR(T, Field) PYBIND11_FIELD_DESCRIPTOR_EX(T, Field, #Field)
# define PYBIND11_FIELD_DESCRIPTOR(T, Field) \
PYBIND11_FIELD_DESCRIPTOR_EX_IMPL((T), Field, #Field)
// The main idea of this macro is borrowed from https://github.com/swansontec/map-macro
// (C) William Swanson, Paul Fultz
@@ -1846,7 +1859,7 @@ private:
# define PYBIND11_NUMPY_DTYPE(Type, ...) \
::pybind11::detail::npy_format_descriptor<Type>::register_dtype( \
::std::vector<::pybind11::detail::field_descriptor>{ \
PYBIND11_MAP_LIST(PYBIND11_FIELD_DESCRIPTOR, Type, __VA_ARGS__)})
PYBIND11_MAP_LIST(PYBIND11_FIELD_DESCRIPTOR_IMPL, (Type), __VA_ARGS__)})
# if defined(_MSC_VER) && !defined(__clang__)
# define PYBIND11_MAP2_LIST_NEXT1(test, next) \
@@ -1868,7 +1881,7 @@ private:
# define PYBIND11_NUMPY_DTYPE_EX(Type, ...) \
::pybind11::detail::npy_format_descriptor<Type>::register_dtype( \
::std::vector<::pybind11::detail::field_descriptor>{ \
PYBIND11_MAP2_LIST(PYBIND11_FIELD_DESCRIPTOR_EX, Type, __VA_ARGS__)})
PYBIND11_MAP2_LIST(PYBIND11_FIELD_DESCRIPTOR_EX_IMPL, (Type), __VA_ARGS__)})
#endif // __CLION_IDE__
+27
View File
@@ -102,6 +102,12 @@ PYBIND11_PACKED(struct StructWithUglyNames {
uint64_t __y__;
});
template <typename T1, typename T2>
struct TemplatedStruct {
T1 a;
T2 b;
};
enum class E1 : int64_t { A = -1, B = 1 };
enum E2 : uint8_t { X = 1, Y = 2 };
@@ -352,6 +358,27 @@ TEST_SUBMODULE(numpy_dtypes, m) {
PYBIND11_NUMPY_DTYPE(EnumStruct, e1, e2);
PYBIND11_NUMPY_DTYPE(ComplexStruct, cflt, cdbl);
// test_templated_dtype
PYBIND11_NUMPY_DTYPE(PYBIND11_TYPE(TemplatedStruct<int32_t, float>), a, b);
PYBIND11_NUMPY_DTYPE_EX(PYBIND11_TYPE(TemplatedStruct<int16_t, uint16_t>), a, "x", b, "y");
m.def("templated_dtypes", []() {
return py::make_tuple(py::dtype::of<TemplatedStruct<int32_t, float>>(),
py::dtype::of<TemplatedStruct<int16_t, uint16_t>>());
});
// test_direct_field_descriptor
m.def("direct_field_descriptors", []() {
py::detail::field_descriptor direct[]
= {PYBIND11_FIELD_DESCRIPTOR(SimpleStruct, uint_),
PYBIND11_FIELD_DESCRIPTOR_EX(SimpleStruct, float_, "flt"),
PYBIND11_FIELD_DESCRIPTOR(PYBIND11_TYPE(TemplatedStruct<int32_t, float>), b)};
py::list names;
for (const auto &fd : direct) {
names.append(fd.name);
}
return names;
});
// ... or after
py::class_<PackedStruct>(m, "PackedStruct");
+11
View File
@@ -205,6 +205,17 @@ def test_dtype(simple_dtype):
assert (m.test_dtype_switch(arr.astype("longdouble")) == arr + 1).all()
def test_templated_dtype():
"""A type spelled with a comma needs PYBIND11_TYPE here."""
plain, renamed = m.templated_dtypes()
assert plain == np.dtype([("a", "i4"), ("b", "f4")])
assert renamed == np.dtype([("x", "i2"), ("y", "u2")])
def test_direct_field_descriptor():
assert m.direct_field_descriptors() == ["uint_", "flt", "b"]
def test_recarray(simple_dtype, packed_dtype):
elements = [(False, 0, 0.0, -0.0), (True, 1, 1.5, -2.5), (False, 2, 3.0, -5.0)]