Deduplicate the BJData ndarray writer's dtype dispatch and drop <map>

write_bjdata_ndarray() built a 12-entry std::map<string_t, CharType> on
every call just to translate the _ArrayType_ name to a dtype marker
(the only reason binary_writer.hpp included <map>), then mapped dtype
to C++ type twice more: once as a switch for the range-check pass and
once as a separate if/else chain for the write pass, with nothing
checking that the two agreed. The caller also ran three at() lookups,
and the callee called value.at(key) about ten more times for the same
three members.

Replace the map with bjdata_ndarray_type_marker(), a plain string
comparison chain (a C++11 constexpr function cannot contain a switch,
so this mirrors binary_reader's own static table style). Replace the
switch/if-chain pair with one write_bjdata_ndarray_elements() that
switches on dtype once and calls a per-type helper -
write_bjdata_ndarray_element<T>() for the eight integer dtypes and
write_bjdata_ndarray_float_element() for 'd' - with a dry_run flag
selecting the range check or the actual write, so the two passes can
no longer disagree on the type. _ArrayType_, _ArraySize_ and
_ArrayData_ are now looked up once into references, and the four
header marker bytes ('[', '$', '#') are written through to_char_type()
like the rest of the UBJSON/BJData writer.

The 'd' (single-precision) rule is left exactly as before, since #5707
is expected to change it separately.

Verified byte-for-byte identical output before/after for every dtype
(including the Draft 2/Draft 3 'byte' fallback and the use_count/
use_type combinations) via a standalone probe, plus round-tripping
through from_bjdata().

Overlaps #5707, which is expected to touch the 'd' dtype case, and
#5518, which is expected to move the write_bjdata_ndarray() call site.

#5710 item 4

Signed-off-by: Niels Lohmann <mail@nlohmann.me>
This commit is contained in:
Niels Lohmann
2026-09-30 18:25:11 +02:00
parent edd18d071e
commit 1841b4f671
2 changed files with 344 additions and 278 deletions
+172 -139
View File
@@ -10,7 +10,6 @@
#include <algorithm> // reverse
#include <array> // array
#include <map> // map
#include <cmath> // isnan, isinf
#include <cstdint> // uint8_t, uint16_t, uint32_t, uint64_t
#include <cstring> // memcpy
@@ -1625,35 +1624,184 @@ class binary_writer
: value_in_range_of<TargetType>(el.template get<std::int64_t>());
}
/*!
@brief look up the BJData ND-array dtype marker for an `_ArrayType_` name
@return the one-character marker, or '\0' if @a name does not name a known dtype
A C++11 `constexpr` function cannot contain a `switch`, so this is a plain
comparison chain instead; it is only reached once per ND-array candidate
object. Keep in sync with binary_reader's `bjd_type_name()`, which maps
the other way.
*/
static CharType bjdata_ndarray_type_marker(const string_t& name)
{
if (name == "uint8")
{
return 'U';
}
if (name == "int8")
{
return 'i';
}
if (name == "uint16")
{
return 'u';
}
if (name == "int16")
{
return 'I';
}
if (name == "uint32")
{
return 'm';
}
if (name == "int32")
{
return 'l';
}
if (name == "uint64")
{
return 'M';
}
if (name == "int64")
{
return 'L';
}
if (name == "single")
{
return 'd';
}
if (name == "double")
{
return 'D';
}
if (name == "char")
{
return 'C';
}
if (name == "byte")
{
return 'B';
}
return '\0';
}
/*!
@brief validate (dry_run) or write one BJData ND-array element of integer dtype @a T
@return whether @a el's value is in range of @a T; always true when @a dry_run is false
*/
template<typename T>
bool write_bjdata_ndarray_element(const BasicJsonType& el, const bool dry_run)
{
if (dry_run)
{
return bjdata_ndarray_value_in_range<T>(el);
}
using storage_type = typename std::conditional<std::is_unsigned<T>::value, std::uint64_t, std::int64_t>::type;
write_number(static_cast<T>(el.template get<storage_type>()), true);
return true;
}
/*!
@brief validate (dry_run) or write one BJData ND-array element of dtype 'd' (single precision)
@return whether @a el's value fits a float without overflow; always true when @a dry_run is false
*/
bool write_bjdata_ndarray_float_element(const BasicJsonType& el, const bool dry_run)
{
const auto dval = el.template get<double>();
if (dry_run)
{
return !std::isfinite(dval) ||
(dval >= static_cast<double>(std::numeric_limits<float>::lowest()) &&
dval <= static_cast<double>((std::numeric_limits<float>::max)()));
}
write_number(static_cast<float>(dval), true);
return true;
}
/*!
@brief validate or write every element of a BJData ND-array's `_ArrayData_`
@param[in] array_data the `_ArrayData_` array
@param[in] dtype the ND-array dtype marker, as returned by bjdata_ndarray_type_marker()
@param[in] dry_run true to only range-check each element, false to write it
@return whether every element is in range for @a dtype (always true when @a dry_run is false)
*/
bool write_bjdata_ndarray_elements(const BasicJsonType& array_data, const CharType dtype, const bool dry_run)
{
for (const auto& el : array_data)
{
bool ok = true;
switch (dtype)
{
case 'U':
case 'C':
case 'B':
ok = write_bjdata_ndarray_element<std::uint8_t>(el, dry_run);
break;
case 'i':
ok = write_bjdata_ndarray_element<std::int8_t>(el, dry_run);
break;
case 'u':
ok = write_bjdata_ndarray_element<std::uint16_t>(el, dry_run);
break;
case 'I':
ok = write_bjdata_ndarray_element<std::int16_t>(el, dry_run);
break;
case 'm':
ok = write_bjdata_ndarray_element<std::uint32_t>(el, dry_run);
break;
case 'l':
ok = write_bjdata_ndarray_element<std::int32_t>(el, dry_run);
break;
case 'M':
ok = write_bjdata_ndarray_element<std::uint64_t>(el, dry_run);
break;
case 'L':
ok = write_bjdata_ndarray_element<std::int64_t>(el, dry_run);
break;
case 'd':
ok = write_bjdata_ndarray_float_element(el, dry_run);
break;
case 'D':
default:
// 'D' (double) already spans the full range of number_float_t
if (!dry_run)
{
write_number(el.template get<double>(), true);
}
break;
}
if (!ok)
{
return false;
}
}
return true;
}
/*!
@return false if the object is successfully converted to a bjdata ndarray, true if the type or size is invalid
*/
bool write_bjdata_ndarray(const typename BasicJsonType::object_t& value, const bool use_count, const bool use_type, const bjdata_version_t bjdata_version)
{
std::map<string_t, CharType> bjdtype = {{"uint8", 'U'}, {"int8", 'i'}, {"uint16", 'u'}, {"int16", 'I'},
{"uint32", 'm'}, {"int32", 'l'}, {"uint64", 'M'}, {"int64", 'L'}, {"single", 'd'}, {"double", 'D'},
{"char", 'C'}, {"byte", 'B'}
};
string_t key = "_ArrayType_";
const auto& array_type = value.at("_ArrayType_");
// the type name is looked up as a string below; a non-string
// annotation (e.g. a number, null, or an array) cannot name a known
// dtype, so it is treated the same as an unrecognized type name and
// falls back to a plain object encoding instead of throwing
// type_error.302 out of get<string_t>()
if (!value.at(key).is_string())
if (!array_type.is_string())
{
return true;
}
// use get<string_t>() instead of static_cast<string_t> to avoid an
// ambiguous conversion under explicit instantiation on C++17 (see #4825)
auto it = bjdtype.find(value.at(key).template get<string_t>());
if (it == bjdtype.end())
const CharType dtype = bjdata_ndarray_type_marker(array_type.template get<string_t>());
if (dtype == '\0')
{
return true;
}
CharType dtype = it->second;
// the 'B' (byte) marker is only defined from BJData Draft 3 onward;
// emitting it under an earlier draft would produce a stream that an
@@ -1665,12 +1813,12 @@ class binary_writer
return true;
}
key = "_ArraySize_";
const auto& array_size = value.at("_ArraySize_");
// the dimensions are written verbatim as the header length below, so a
// value that is not an array cannot produce a valid one: null emits 'Z'
// and an object emits '{', neither of which a reader accepts after '#'.
// Such an object is not a valid ndarray and falls back to a plain object.
if (!value.at(key).is_array())
if (!array_size.is_array())
{
return true;
}
@@ -1680,7 +1828,7 @@ class binary_writer
// dimension, or a 1xN row vector is read back as a plain array, which
// would silently drop the annotation, so such an object falls back to
// a plain object encoding instead
const auto& dims = value.at(key);
const auto& dims = array_size;
if (dims.size() < 2 || (dims.size() == 2 && dims.at(0).is_number_integer() && dims.at(0).template get<std::int64_t>() == 1))
{
return true;
@@ -1727,8 +1875,8 @@ class binary_writer
// to be an array: size() is 0 for null and 1 for any other scalar, and
// iterating an object visits its values, so any of these could match
// the dimensions by accident and be encoded as an unrelated ND-array
key = "_ArrayData_";
if (!value.at(key).is_array() || value.at(key).size() != len)
const auto& array_data = value.at("_ArrayData_");
if (!array_data.is_array() || array_data.size() != len)
{
return true;
}
@@ -1743,7 +1891,7 @@ class binary_writer
// API stores int literals as signed), so both are accepted here and the
// writes below go through get<>, which reads the member that is active.
const bool ndarray_is_float = (dtype == 'd' || dtype == 'D');
for (const auto& el : value.at(key))
for (const auto& el : array_data)
{
if (ndarray_is_float ? !el.is_number_float() : !el.is_number_integer())
{
@@ -1756,134 +1904,19 @@ class binary_writer
// wrap (integers) or overflow to infinity (the "single" precision
// float) instead of being reported, so such an object falls back to
// a plain object encoding as well
for (const auto& el : value.at(key))
if (!write_bjdata_ndarray_elements(array_data, dtype, true))
{
bool in_range = true;
switch (dtype)
{
case 'U':
case 'C':
case 'B':
in_range = bjdata_ndarray_value_in_range<std::uint8_t>(el);
break;
case 'i':
in_range = bjdata_ndarray_value_in_range<std::int8_t>(el);
break;
case 'u':
in_range = bjdata_ndarray_value_in_range<std::uint16_t>(el);
break;
case 'I':
in_range = bjdata_ndarray_value_in_range<std::int16_t>(el);
break;
case 'm':
in_range = bjdata_ndarray_value_in_range<std::uint32_t>(el);
break;
case 'l':
in_range = bjdata_ndarray_value_in_range<std::int32_t>(el);
break;
case 'M':
in_range = bjdata_ndarray_value_in_range<std::uint64_t>(el);
break;
case 'L':
in_range = bjdata_ndarray_value_in_range<std::int64_t>(el);
break;
case 'd':
{
const auto dval = el.template get<double>();
in_range = !std::isfinite(dval) ||
(dval >= static_cast<double>(std::numeric_limits<float>::lowest()) &&
dval <= static_cast<double>((std::numeric_limits<float>::max)()));
break;
}
default:
// 'D' (double) already spans the full range of number_float_t
break;
}
if (!in_range)
{
return true;
}
return true;
}
oa.write_character('[');
oa.write_character('$');
oa.write_character(to_char_type('['));
oa.write_character(to_char_type('$'));
oa.write_character(dtype);
oa.write_character('#');
oa.write_character(to_char_type('#'));
key = "_ArraySize_";
write_ubjson(value.at(key), use_count, use_type, true, true, bjdata_version);
write_ubjson(array_size, use_count, use_type, true, true, bjdata_version);
key = "_ArrayData_";
if (dtype == 'U' || dtype == 'C' || dtype == 'B')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::uint8_t>(el.template get<std::uint64_t>()), true);
}
}
else if (dtype == 'i')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::int8_t>(el.template get<std::int64_t>()), true);
}
}
else if (dtype == 'u')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::uint16_t>(el.template get<std::uint64_t>()), true);
}
}
else if (dtype == 'I')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::int16_t>(el.template get<std::int64_t>()), true);
}
}
else if (dtype == 'm')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::uint32_t>(el.template get<std::uint64_t>()), true);
}
}
else if (dtype == 'l')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::int32_t>(el.template get<std::int64_t>()), true);
}
}
else if (dtype == 'M')
{
for (const auto& el : value.at(key))
{
write_number(el.template get<std::uint64_t>(), true);
}
}
else if (dtype == 'L')
{
for (const auto& el : value.at(key))
{
write_number(el.template get<std::int64_t>(), true);
}
}
else if (dtype == 'd')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<float>(el.template get<double>()), true);
}
}
else if (dtype == 'D')
{
for (const auto& el : value.at(key))
{
write_number(el.template get<double>(), true);
}
}
write_bjdata_ndarray_elements(array_data, dtype, false);
return false;
}
+172 -139
View File
@@ -20078,7 +20078,6 @@ NLOHMANN_JSON_NAMESPACE_END
#include <algorithm> // reverse
#include <array> // array
#include <map> // map
#include <cmath> // isnan, isinf
#include <cstdint> // uint8_t, uint16_t, uint32_t, uint64_t
#include <cstring> // memcpy
@@ -21918,35 +21917,184 @@ class binary_writer
: value_in_range_of<TargetType>(el.template get<std::int64_t>());
}
/*!
@brief look up the BJData ND-array dtype marker for an `_ArrayType_` name
@return the one-character marker, or '\0' if @a name does not name a known dtype
A C++11 `constexpr` function cannot contain a `switch`, so this is a plain
comparison chain instead; it is only reached once per ND-array candidate
object. Keep in sync with binary_reader's `bjd_type_name()`, which maps
the other way.
*/
static CharType bjdata_ndarray_type_marker(const string_t& name)
{
if (name == "uint8")
{
return 'U';
}
if (name == "int8")
{
return 'i';
}
if (name == "uint16")
{
return 'u';
}
if (name == "int16")
{
return 'I';
}
if (name == "uint32")
{
return 'm';
}
if (name == "int32")
{
return 'l';
}
if (name == "uint64")
{
return 'M';
}
if (name == "int64")
{
return 'L';
}
if (name == "single")
{
return 'd';
}
if (name == "double")
{
return 'D';
}
if (name == "char")
{
return 'C';
}
if (name == "byte")
{
return 'B';
}
return '\0';
}
/*!
@brief validate (dry_run) or write one BJData ND-array element of integer dtype @a T
@return whether @a el's value is in range of @a T; always true when @a dry_run is false
*/
template<typename T>
bool write_bjdata_ndarray_element(const BasicJsonType& el, const bool dry_run)
{
if (dry_run)
{
return bjdata_ndarray_value_in_range<T>(el);
}
using storage_type = typename std::conditional<std::is_unsigned<T>::value, std::uint64_t, std::int64_t>::type;
write_number(static_cast<T>(el.template get<storage_type>()), true);
return true;
}
/*!
@brief validate (dry_run) or write one BJData ND-array element of dtype 'd' (single precision)
@return whether @a el's value fits a float without overflow; always true when @a dry_run is false
*/
bool write_bjdata_ndarray_float_element(const BasicJsonType& el, const bool dry_run)
{
const auto dval = el.template get<double>();
if (dry_run)
{
return !std::isfinite(dval) ||
(dval >= static_cast<double>(std::numeric_limits<float>::lowest()) &&
dval <= static_cast<double>((std::numeric_limits<float>::max)()));
}
write_number(static_cast<float>(dval), true);
return true;
}
/*!
@brief validate or write every element of a BJData ND-array's `_ArrayData_`
@param[in] array_data the `_ArrayData_` array
@param[in] dtype the ND-array dtype marker, as returned by bjdata_ndarray_type_marker()
@param[in] dry_run true to only range-check each element, false to write it
@return whether every element is in range for @a dtype (always true when @a dry_run is false)
*/
bool write_bjdata_ndarray_elements(const BasicJsonType& array_data, const CharType dtype, const bool dry_run)
{
for (const auto& el : array_data)
{
bool ok = true;
switch (dtype)
{
case 'U':
case 'C':
case 'B':
ok = write_bjdata_ndarray_element<std::uint8_t>(el, dry_run);
break;
case 'i':
ok = write_bjdata_ndarray_element<std::int8_t>(el, dry_run);
break;
case 'u':
ok = write_bjdata_ndarray_element<std::uint16_t>(el, dry_run);
break;
case 'I':
ok = write_bjdata_ndarray_element<std::int16_t>(el, dry_run);
break;
case 'm':
ok = write_bjdata_ndarray_element<std::uint32_t>(el, dry_run);
break;
case 'l':
ok = write_bjdata_ndarray_element<std::int32_t>(el, dry_run);
break;
case 'M':
ok = write_bjdata_ndarray_element<std::uint64_t>(el, dry_run);
break;
case 'L':
ok = write_bjdata_ndarray_element<std::int64_t>(el, dry_run);
break;
case 'd':
ok = write_bjdata_ndarray_float_element(el, dry_run);
break;
case 'D':
default:
// 'D' (double) already spans the full range of number_float_t
if (!dry_run)
{
write_number(el.template get<double>(), true);
}
break;
}
if (!ok)
{
return false;
}
}
return true;
}
/*!
@return false if the object is successfully converted to a bjdata ndarray, true if the type or size is invalid
*/
bool write_bjdata_ndarray(const typename BasicJsonType::object_t& value, const bool use_count, const bool use_type, const bjdata_version_t bjdata_version)
{
std::map<string_t, CharType> bjdtype = {{"uint8", 'U'}, {"int8", 'i'}, {"uint16", 'u'}, {"int16", 'I'},
{"uint32", 'm'}, {"int32", 'l'}, {"uint64", 'M'}, {"int64", 'L'}, {"single", 'd'}, {"double", 'D'},
{"char", 'C'}, {"byte", 'B'}
};
string_t key = "_ArrayType_";
const auto& array_type = value.at("_ArrayType_");
// the type name is looked up as a string below; a non-string
// annotation (e.g. a number, null, or an array) cannot name a known
// dtype, so it is treated the same as an unrecognized type name and
// falls back to a plain object encoding instead of throwing
// type_error.302 out of get<string_t>()
if (!value.at(key).is_string())
if (!array_type.is_string())
{
return true;
}
// use get<string_t>() instead of static_cast<string_t> to avoid an
// ambiguous conversion under explicit instantiation on C++17 (see #4825)
auto it = bjdtype.find(value.at(key).template get<string_t>());
if (it == bjdtype.end())
const CharType dtype = bjdata_ndarray_type_marker(array_type.template get<string_t>());
if (dtype == '\0')
{
return true;
}
CharType dtype = it->second;
// the 'B' (byte) marker is only defined from BJData Draft 3 onward;
// emitting it under an earlier draft would produce a stream that an
@@ -21958,12 +22106,12 @@ class binary_writer
return true;
}
key = "_ArraySize_";
const auto& array_size = value.at("_ArraySize_");
// the dimensions are written verbatim as the header length below, so a
// value that is not an array cannot produce a valid one: null emits 'Z'
// and an object emits '{', neither of which a reader accepts after '#'.
// Such an object is not a valid ndarray and falls back to a plain object.
if (!value.at(key).is_array())
if (!array_size.is_array())
{
return true;
}
@@ -21973,7 +22121,7 @@ class binary_writer
// dimension, or a 1xN row vector is read back as a plain array, which
// would silently drop the annotation, so such an object falls back to
// a plain object encoding instead
const auto& dims = value.at(key);
const auto& dims = array_size;
if (dims.size() < 2 || (dims.size() == 2 && dims.at(0).is_number_integer() && dims.at(0).template get<std::int64_t>() == 1))
{
return true;
@@ -22020,8 +22168,8 @@ class binary_writer
// to be an array: size() is 0 for null and 1 for any other scalar, and
// iterating an object visits its values, so any of these could match
// the dimensions by accident and be encoded as an unrelated ND-array
key = "_ArrayData_";
if (!value.at(key).is_array() || value.at(key).size() != len)
const auto& array_data = value.at("_ArrayData_");
if (!array_data.is_array() || array_data.size() != len)
{
return true;
}
@@ -22036,7 +22184,7 @@ class binary_writer
// API stores int literals as signed), so both are accepted here and the
// writes below go through get<>, which reads the member that is active.
const bool ndarray_is_float = (dtype == 'd' || dtype == 'D');
for (const auto& el : value.at(key))
for (const auto& el : array_data)
{
if (ndarray_is_float ? !el.is_number_float() : !el.is_number_integer())
{
@@ -22049,134 +22197,19 @@ class binary_writer
// wrap (integers) or overflow to infinity (the "single" precision
// float) instead of being reported, so such an object falls back to
// a plain object encoding as well
for (const auto& el : value.at(key))
if (!write_bjdata_ndarray_elements(array_data, dtype, true))
{
bool in_range = true;
switch (dtype)
{
case 'U':
case 'C':
case 'B':
in_range = bjdata_ndarray_value_in_range<std::uint8_t>(el);
break;
case 'i':
in_range = bjdata_ndarray_value_in_range<std::int8_t>(el);
break;
case 'u':
in_range = bjdata_ndarray_value_in_range<std::uint16_t>(el);
break;
case 'I':
in_range = bjdata_ndarray_value_in_range<std::int16_t>(el);
break;
case 'm':
in_range = bjdata_ndarray_value_in_range<std::uint32_t>(el);
break;
case 'l':
in_range = bjdata_ndarray_value_in_range<std::int32_t>(el);
break;
case 'M':
in_range = bjdata_ndarray_value_in_range<std::uint64_t>(el);
break;
case 'L':
in_range = bjdata_ndarray_value_in_range<std::int64_t>(el);
break;
case 'd':
{
const auto dval = el.template get<double>();
in_range = !std::isfinite(dval) ||
(dval >= static_cast<double>(std::numeric_limits<float>::lowest()) &&
dval <= static_cast<double>((std::numeric_limits<float>::max)()));
break;
}
default:
// 'D' (double) already spans the full range of number_float_t
break;
}
if (!in_range)
{
return true;
}
return true;
}
oa.write_character('[');
oa.write_character('$');
oa.write_character(to_char_type('['));
oa.write_character(to_char_type('$'));
oa.write_character(dtype);
oa.write_character('#');
oa.write_character(to_char_type('#'));
key = "_ArraySize_";
write_ubjson(value.at(key), use_count, use_type, true, true, bjdata_version);
write_ubjson(array_size, use_count, use_type, true, true, bjdata_version);
key = "_ArrayData_";
if (dtype == 'U' || dtype == 'C' || dtype == 'B')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::uint8_t>(el.template get<std::uint64_t>()), true);
}
}
else if (dtype == 'i')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::int8_t>(el.template get<std::int64_t>()), true);
}
}
else if (dtype == 'u')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::uint16_t>(el.template get<std::uint64_t>()), true);
}
}
else if (dtype == 'I')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::int16_t>(el.template get<std::int64_t>()), true);
}
}
else if (dtype == 'm')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::uint32_t>(el.template get<std::uint64_t>()), true);
}
}
else if (dtype == 'l')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<std::int32_t>(el.template get<std::int64_t>()), true);
}
}
else if (dtype == 'M')
{
for (const auto& el : value.at(key))
{
write_number(el.template get<std::uint64_t>(), true);
}
}
else if (dtype == 'L')
{
for (const auto& el : value.at(key))
{
write_number(el.template get<std::int64_t>(), true);
}
}
else if (dtype == 'd')
{
for (const auto& el : value.at(key))
{
write_number(static_cast<float>(el.template get<double>()), true);
}
}
else if (dtype == 'D')
{
for (const auto& el : value.at(key))
{
write_number(el.template get<double>(), true);
}
}
write_bjdata_ndarray_elements(array_data, dtype, false);
return false;
}