Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 9 additions & 9 deletions msgpack/fallback.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
12 changes: 12 additions & 0 deletions test/test_obj.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading