diff --git a/msgpack/_packer.pyx b/msgpack/_packer.pyx index c44c6366..4ecbb199 100644 --- a/msgpack/_packer.pyx +++ b/msgpack/_packer.pyx @@ -332,13 +332,17 @@ cdef class Packer: (`len(pairs)` and `for k, v in pairs:` should be supported.) """ self._check_exports() - size = len(pairs) - if size > ITEM_LIMIT: - raise ValueError("map too large") - msgpack_pack_map(&self.pk, size) - for k, v in pairs: - self._pack(k) - self._pack(v) + try: + size = len(pairs) + if size > ITEM_LIMIT: + raise ValueError("map too large") + msgpack_pack_map(&self.pk, size) + for k, v in pairs: + self._pack(k) + self._pack(v) + except: + self.pk.length = 0 + raise if self.autoreset: buf = PyBytes_FromStringAndSize(self.pk.buf, self.pk.length) self.pk.length = 0 diff --git a/msgpack/fallback.py b/msgpack/fallback.py index cee48e9a..44baca56 100644 --- a/msgpack/fallback.py +++ b/msgpack/fallback.py @@ -830,7 +830,11 @@ def pack(self, obj): return ret def pack_map_pairs(self, pairs): - self._pack_map_pairs(len(pairs), pairs) + try: + self._pack_map_pairs(len(pairs), pairs) + except: + self._buffer = BytesIO() + raise if self._autoreset: ret = self._buffer.getvalue() self._buffer = BytesIO() diff --git a/test/test_pack.py b/test/test_pack.py index 9ca6e182..d428e953 100644 --- a/test/test_pack.py +++ b/test/test_pack.py @@ -190,6 +190,37 @@ def test_pairlist(): assert pairlist == unpacked +@pytest.mark.parametrize("autoreset", [True, False]) +@pytest.mark.parametrize("method", ["pack", "pack_map_pairs"]) +def test_packer_resets_after_default_error(autoreset, method): + class Invoice: + def __init__(self, ready): + self.ready = ready + + def default(invoice): + if not invoice.ready: + raise ValueError("invoice not ready") + return {"amount": 15} + + packer = Packer(default=default, autoreset=autoreset) + packer.pack({"previous": 1}) + pack = getattr(packer, method) + failed = [("id", 1), ("invoice", Invoice(False))] + with pytest.raises(ValueError, match="invoice not ready"): + pack(dict(failed) if method == "pack" else failed) + assert packer.bytes() == b"" + + valid = [("invoice", Invoice(True))] + packed = pack(dict(valid) if method == "pack" else valid) + if autoreset: + assert unpackb(packed) == {"invoice": {"amount": 15}} + else: + packer.pack({"next": 2}) + unpacker = Unpacker() + unpacker.feed(packer.bytes()) + assert list(unpacker) == [{"invoice": {"amount": 15}}, {"next": 2}] + + def test_get_buffer(): packer = Packer(autoreset=0, use_bin_type=True) packer.pack([1, 2])