Merge branch 'fix-mmap-buf' into madhava/python_310

This commit is contained in:
Madhava Jay
2022-05-24 10:42:57 +10:00
2 changed files with 45 additions and 22 deletions

View File

@@ -22,6 +22,7 @@ from libc.string cimport memcpy
import array import array
import asyncio import asyncio
import collections as _collections import collections as _collections
import contextlib
import enum as _enum import enum as _enum
import inspect as _inspect import inspect as _inspect
import os as _os import os as _os
@@ -3320,6 +3321,7 @@ class _StructModule(object):
reader = _MultipleBytesPackedMessageReader(buf, self.schema, traversal_limit_in_words, nesting_limit) reader = _MultipleBytesPackedMessageReader(buf, self.schema, traversal_limit_in_words, nesting_limit)
return reader return reader
@contextlib.contextmanager
def from_bytes(self, buf, traversal_limit_in_words=None, nesting_limit=None, builder=False): def from_bytes(self, buf, traversal_limit_in_words=None, nesting_limit=None, builder=False):
"""Returns a Reader for the unpacked object in buf. """Returns a Reader for the unpacked object in buf.
@@ -3340,13 +3342,18 @@ class _StructModule(object):
:rtype: :class:`_DynamicStructReader` or :class:`_DynamicStructBuilder` :rtype: :class:`_DynamicStructReader` or :class:`_DynamicStructBuilder`
""" """
message = None
try:
if builder: if builder:
# message = _FlatMessageBuilder(buf) # message = _FlatMessageBuilder(buf)
message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit) message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit)
return message.get_root(self.schema).as_builder() yield message.get_root(self.schema).as_builder()
else: else:
message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit) message = _FlatArrayMessageReader(buf, traversal_limit_in_words, nesting_limit)
return message.get_root(self.schema) yield message.get_root(self.schema)
finally:
if message:
message.close()
def from_segments(self, segments, traversal_limit_in_words=None, nesting_limit=None): def from_segments(self, segments, traversal_limit_in_words=None, nesting_limit=None):
"""Returns a Reader for a list of segment bytes. """Returns a Reader for a list of segment bytes.
@@ -4088,15 +4095,22 @@ cdef class _AlignedBuffer:
cdef class _BufferView: cdef class _BufferView:
cdef Py_buffer view cdef Py_buffer view
cdef char * buf cdef char * buf
cdef int closed
def __init__(self, other): def __init__(self, other):
cdef int ret = PyObject_GetBuffer(other, &self.view, PyBUF_SIMPLE) cdef int ret = PyObject_GetBuffer(other, &self.view, PyBUF_SIMPLE)
if ret < 0: if ret < 0:
raise ValueError("Invalid buffer passed to BufferView") raise ValueError("Invalid buffer passed to BufferView")
self.buf = <char*>self.view.buf self.buf = <char*>self.view.buf
self.closed = False
def close(self):
if not self.closed:
PyBuffer_Release(&self.view)
self.closed = True
def __dealloc__(self): def __dealloc__(self):
PyBuffer_Release(&self.view) self.close()
@cython.internal @cython.internal
@@ -4133,6 +4147,7 @@ cdef class _FlatArrayMessageReaderAligned(_MessageReader):
@cython.internal @cython.internal
cdef class _FlatArrayMessageReader(_MessageReader): cdef class _FlatArrayMessageReader(_MessageReader):
cdef object _object_to_pin cdef object _object_to_pin
cdef _BufferView _buffer_view
def __init__(self, buf, traversal_limit_in_words=None, nesting_limit=None): 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) 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 self._object_to_pin = aligned
else: else:
self._object_to_pin = buf self._object_to_pin = buf
self._buffer_view = None
elif PyObject_CheckBuffer(buf): elif PyObject_CheckBuffer(buf):
view = _BufferView(buf) view = _BufferView(buf)
ptr = view.buf ptr = view.buf
self._object_to_pin = view self._object_to_pin = view
self._buffer_view = view
else: else:
raise TypeError('expected buffer-like object in FlatArrayMessageReader') raise TypeError('expected buffer-like object in FlatArrayMessageReader')
@@ -4162,7 +4179,12 @@ cdef class _FlatArrayMessageReader(_MessageReader):
schema_cpp.WordArrayPtr(<schema_cpp.word*>ptr, sz//8), schema_cpp.WordArrayPtr(<schema_cpp.word*>ptr, sz//8),
opts) opts)
def close(self):
if self._buffer_view:
self._buffer_view.close()
def __dealloc__(self): def __dealloc__(self):
self.close()
del self.thisptr del self.thisptr

View File

@@ -46,7 +46,7 @@ def test_roundtrip_bytes(all_types):
test_regression.init_all_types(msg) test_regression.init_all_types(msg)
message_bytes = msg.to_bytes() message_bytes = msg.to_bytes()
msg = all_types.TestAllTypes.from_bytes(message_bytes) with all_types.TestAllTypes.from_bytes(message_bytes) as msg:
test_regression.check_all_types(msg) test_regression.check_all_types(msg)
@@ -77,7 +77,7 @@ def test_roundtrip_bytes_mmap(all_types):
f.seek(0) f.seek(0)
memory = mmap.mmap(f.fileno(), length) memory = mmap.mmap(f.fileno(), length)
msg = all_types.TestAllTypes.from_bytes(memory) with all_types.TestAllTypes.from_bytes(memory) as msg:
test_regression.check_all_types(msg) test_regression.check_all_types(msg)
@@ -91,7 +91,7 @@ def test_roundtrip_bytes_buffer(all_types):
b = msg.to_bytes() b = msg.to_bytes()
v = memoryview(b) v = memoryview(b)
try: try:
msg = all_types.TestAllTypes.from_bytes(v) with all_types.TestAllTypes.from_bytes(v) as msg:
test_regression.check_all_types(msg) test_regression.check_all_types(msg)
finally: finally:
v.release() v.release()
@@ -99,7 +99,8 @@ def test_roundtrip_bytes_buffer(all_types):
def test_roundtrip_bytes_fail(all_types): def test_roundtrip_bytes_fail(all_types):
with pytest.raises(TypeError): with pytest.raises(TypeError):
all_types.TestAllTypes.from_bytes(42) with all_types.TestAllTypes.from_bytes(42) as msg:
pass
@pytest.mark.skipif( @pytest.mark.skipif(
@@ -232,12 +233,12 @@ def test_from_bytes_traversal_limit(all_types):
bld.init("structList", size) bld.init("structList", size)
data = bld.to_bytes() data = bld.to_bytes()
msg = all_types.TestAllTypes.from_bytes(data) with all_types.TestAllTypes.from_bytes(data) as msg:
with pytest.raises(capnp.KjException): with pytest.raises(capnp.KjException):
for i in range(0, size): for i in range(0, size):
msg.structList[i].uInt8Field == 0 msg.structList[i].uInt8Field == 0
msg = all_types.TestAllTypes.from_bytes(data, traversal_limit_in_words=2**62) with all_types.TestAllTypes.from_bytes(data, traversal_limit_in_words=2 ** 62) as msg:
for i in range(0, size): for i in range(0, size):
assert msg.structList[i].uInt8Field == 0 assert msg.structList[i].uInt8Field == 0