diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index dd0f9b8..e277ad9 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -22,6 +22,7 @@ from libc.string cimport memcpy import array import asyncio import collections as _collections +import contextlib import enum as _enum import inspect as _inspect import os as _os @@ -3320,6 +3321,7 @@ class _StructModule(object): reader = _MultipleBytesPackedMessageReader(buf, self.schema, traversal_limit_in_words, nesting_limit) return reader + @contextlib.contextmanager def from_bytes(self, buf, traversal_limit_in_words=None, nesting_limit=None, builder=False): """Returns a Reader for the unpacked object in buf. @@ -3340,13 +3342,18 @@ class _StructModule(object): :rtype: :class:`_DynamicStructReader` or :class:`_DynamicStructBuilder` """ - if builder: - # message = _FlatMessageBuilder(buf) - message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit) - return message.get_root(self.schema).as_builder() - else: - message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit) - return message.get_root(self.schema) + message = None + try: + if builder: + # message = _FlatMessageBuilder(buf) + message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit) + yield message.get_root(self.schema).as_builder() + else: + message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit) + yield message.get_root(self.schema) + finally: + if message: + message.close() def from_segments(self, segments, traversal_limit_in_words=None, nesting_limit=None): """Returns a Reader for a list of segment bytes. @@ -4088,15 +4095,22 @@ cdef class _AlignedBuffer: cdef class _BufferView: cdef Py_buffer view cdef char * buf + cdef int closed def __init__(self, other): cdef int ret = PyObject_GetBuffer(other, &self.view, PyBUF_SIMPLE) if ret < 0: raise ValueError("Invalid buffer passed to BufferView") self.buf = self.view.buf + self.closed = False + + def close(self): + if not self.closed: + PyBuffer_Release(&self.view) + self.closed = True def __dealloc__(self): - PyBuffer_Release(&self.view) + self.close() @cython.internal @@ -4133,6 +4147,7 @@ cdef class _FlatArrayMessageReaderAligned(_MessageReader): @cython.internal cdef class _FlatArrayMessageReader(_MessageReader): cdef object _object_to_pin + cdef _BufferView _buffer_view def __init__(self, buf, traversal_limit_in_words=None, nesting_limit=None): cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit) @@ -4151,10 +4166,12 @@ cdef class _FlatArrayMessageReader(_MessageReader): self._object_to_pin = aligned else: self._object_to_pin = buf + self._buffer_view = None elif PyObject_CheckBuffer(buf): view = _BufferView(buf) ptr = view.buf self._object_to_pin = view + self._buffer_view = view else: raise TypeError('expected buffer-like object in FlatArrayMessageReader') @@ -4162,7 +4179,12 @@ cdef class _FlatArrayMessageReader(_MessageReader): schema_cpp.WordArrayPtr(ptr, sz//8), opts) + def close(self): + if self._buffer_view: + self._buffer_view.close() + def __dealloc__(self): + self.close() del self.thisptr diff --git a/test/test_serialization.py b/test/test_serialization.py index 558e3c1..09e1278 100644 --- a/test/test_serialization.py +++ b/test/test_serialization.py @@ -46,8 +46,8 @@ def test_roundtrip_bytes(all_types): test_regression.init_all_types(msg) message_bytes = msg.to_bytes() - msg = all_types.TestAllTypes.from_bytes(message_bytes) - test_regression.check_all_types(msg) + with all_types.TestAllTypes.from_bytes(message_bytes) as msg: + test_regression.check_all_types(msg) @pytest.mark.skipif( @@ -77,8 +77,8 @@ def test_roundtrip_bytes_mmap(all_types): f.seek(0) memory = mmap.mmap(f.fileno(), length) - msg = all_types.TestAllTypes.from_bytes(memory) - test_regression.check_all_types(msg) + with all_types.TestAllTypes.from_bytes(memory) as msg: + test_regression.check_all_types(msg) @pytest.mark.skipif( @@ -91,15 +91,16 @@ def test_roundtrip_bytes_buffer(all_types): b = msg.to_bytes() v = memoryview(b) try: - msg = all_types.TestAllTypes.from_bytes(v) - test_regression.check_all_types(msg) + with all_types.TestAllTypes.from_bytes(v) as msg: + test_regression.check_all_types(msg) finally: v.release() def test_roundtrip_bytes_fail(all_types): with pytest.raises(TypeError): - all_types.TestAllTypes.from_bytes(42) + with all_types.TestAllTypes.from_bytes(42) as msg: + pass @pytest.mark.skipif( @@ -232,14 +233,14 @@ def test_from_bytes_traversal_limit(all_types): bld.init("structList", size) data = bld.to_bytes() - msg = all_types.TestAllTypes.from_bytes(data) - with pytest.raises(capnp.KjException): - for i in range(0, size): - msg.structList[i].uInt8Field == 0 + with all_types.TestAllTypes.from_bytes(data) as msg: + with pytest.raises(capnp.KjException): + for i in range(0, size): + msg.structList[i].uInt8Field == 0 - msg = all_types.TestAllTypes.from_bytes(data, traversal_limit_in_words=2**62) - for i in range(0, size): - assert msg.structList[i].uInt8Field == 0 + with all_types.TestAllTypes.from_bytes(data, traversal_limit_in_words=2 ** 62) as msg: + for i in range(0, size): + assert msg.structList[i].uInt8Field == 0 def test_from_bytes_packed_traversal_limit(all_types):