diff --git a/msgpack/fallback.py b/msgpack/fallback.py index 569650af..cee48e9a 100644 --- a/msgpack/fallback.py +++ b/msgpack/fallback.py @@ -529,15 +529,15 @@ def _unpack(self, execute=EX_CONSTRUCT): self._unpack(EX_SKIP) return if self._object_pairs_hook is not None: - - def _gen(): - for _ in range(n): - key = self._unpack(EX_CONSTRUCT) - if self._strict_map_key and type(key) not in (str, bytes): - raise ValueError("%s is not allowed for map key" % str(type(key))) - yield key, self._unpack(EX_CONSTRUCT) - - ret = self._object_pairs_hook(_gen()) + # Pass a list, as the C extension does, so the whole map is + # consumed even if the hook does not iterate it. + pairs = [] + for _ in range(n): + key = self._unpack(EX_CONSTRUCT) + if self._strict_map_key and type(key) not in (str, bytes): + raise ValueError("%s is not allowed for map key" % str(type(key))) + pairs.append((key, self._unpack(EX_CONSTRUCT))) + ret = self._object_pairs_hook(pairs) else: ret = {} for _ in range(n): diff --git a/test/test_obj.py b/test/test_obj.py index 23be06d5..1866bc67 100644 --- a/test/test_obj.py +++ b/test/test_obj.py @@ -41,6 +41,18 @@ def test_decode_pairs_hook(): assert unpacked[1] == prod_sum +def test_decode_pairs_hook_receives_list(): + def reject_duplicate_keys(pairs): + keys = [k for k, _ in pairs] + assert len(keys) == len(set(keys)) + return dict(pairs) + + packed = packb([{"a": 1, "b": 2}, 3]) + assert unpackb(packed, object_pairs_hook=reject_duplicate_keys) == [{"a": 1, "b": 2}, 3] + assert unpackb(packed, object_pairs_hook=lambda pairs: pairs[0]) == [("a", 1), 3] + assert unpackb(packed, object_pairs_hook=lambda pairs: None) == [None, 3] + + def test_only_one_obj_hook(): with raises(TypeError): unpackb(b"", object_hook=lambda x: x, object_pairs_hook=lambda x: x)