Merge pull request #126 from asilversempirical/pass_reader_opts_clean

Let TwoParty Clients and Servers take ReaderOptions
This commit is contained in:
Jason Paryani
2017-04-09 16:52:36 -07:00
committed by GitHub
3 changed files with 46 additions and 54 deletions

View File

@@ -4,7 +4,7 @@
cdef extern from "capnp/helpers/checkCompiler.h": cdef extern from "capnp/helpers/checkCompiler.h":
pass pass
from schema_cpp cimport Node, Data, StructNode, EnumNode, InterfaceNode, MessageBuilder, MessageReader from schema_cpp cimport Node, Data, StructNode, EnumNode, InterfaceNode, MessageBuilder, MessageReader, ReaderOptions
from capnp.helpers.non_circular cimport PythonInterfaceDynamicImpl, reraise_kj_exception, PyRefCounter, PyRestorer, PyEventPort, ErrorHandler from capnp.helpers.non_circular cimport PythonInterfaceDynamicImpl, reraise_kj_exception, PyRefCounter, PyRestorer, PyEventPort, ErrorHandler
from capnp.includes.types cimport * from capnp.includes.types cimport *
@@ -45,7 +45,7 @@ cdef extern from "kj/exception.h" namespace " ::kj":
cdef extern from "kj/memory.h" namespace " ::kj": cdef extern from "kj/memory.h" namespace " ::kj":
cdef cppclass Own[T]: cdef cppclass Own[T]:
T& operator*() T& operator*()
Own[TwoPartyVatNetwork] makeTwoPartyVatNetwork" ::kj::heap< ::capnp::TwoPartyVatNetwork>"(AsyncIoStream& stream, Side) Own[TwoPartyVatNetwork] makeTwoPartyVatNetwork" ::kj::heap< ::capnp::TwoPartyVatNetwork>"(AsyncIoStream& stream, Side, ReaderOptions)
Own[PromiseFulfillerPair] copyPromiseFulfillerPair" ::kj::heap< ::kj::PromiseFulfillerPair<void> >"(PromiseFulfillerPair&) Own[PromiseFulfillerPair] copyPromiseFulfillerPair" ::kj::heap< ::kj::PromiseFulfillerPair<void> >"(PromiseFulfillerPair&)
Own[PyRefCounter] makePyRefCounter" ::kj::heap< PyRefCounter >"(PyObject *) Own[PyRefCounter] makePyRefCounter" ::kj::heap< PyRefCounter >"(PyObject *)
@@ -341,7 +341,7 @@ cdef extern from "capnp/rpc-twoparty.h" namespace " ::capnp":
cdef Side SERVER" ::capnp::rpc::twoparty::Side::SERVER" cdef Side SERVER" ::capnp::rpc::twoparty::Side::SERVER"
cdef cppclass TwoPartyVatNetwork: cdef cppclass TwoPartyVatNetwork:
TwoPartyVatNetwork(EventLoop &, AsyncIoStream& stream, Side) TwoPartyVatNetwork(EventLoop &, AsyncIoStream& stream, Side, ReaderOptions)
VoidPromise onDisconnect() VoidPromise onDisconnect()
VoidPromise onDrained() VoidPromise onDrained()
RpcSystem makeRpcServer(TwoPartyVatNetwork&, PyRestorer&) RpcSystem makeRpcServer(TwoPartyVatNetwork&, PyRestorer&)

View File

@@ -264,6 +264,14 @@ cdef public object get_exception_info(object exc_type, object exc_obj, object ex
except: except:
return (b'', 0, b"Couldn't determine python exception") return (b'', 0, b"Couldn't determine python exception")
cdef schema_cpp.ReaderOptions make_reader_opts(traversal_limit_in_words, nesting_limit) with gil:
cdef schema_cpp.ReaderOptions opts
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
return opts
ctypedef fused _DynamicStructReaderOrBuilder: ctypedef fused _DynamicStructReaderOrBuilder:
_DynamicStructReader _DynamicStructReader
_DynamicStructBuilder _DynamicStructBuilder
@@ -2217,9 +2225,9 @@ cdef class _TwoPartyVatNetwork:
cdef Own[C_TwoPartyVatNetwork] thisptr cdef Own[C_TwoPartyVatNetwork] thisptr
cdef _AsyncIoStream stream cdef _AsyncIoStream stream
cdef _init(self, _AsyncIoStream stream, Side side): cdef _init(self, _AsyncIoStream stream, Side side, schema_cpp.ReaderOptions opts):
self.stream = stream self.stream = stream
self.thisptr = makeTwoPartyVatNetwork(deref(stream.thisptr), side) self.thisptr = makeTwoPartyVatNetwork(deref(stream.thisptr), side, opts)
return self return self
cpdef on_disconnect(self) except +reraise_kj_exception: cpdef on_disconnect(self) except +reraise_kj_exception:
@@ -2244,13 +2252,15 @@ cdef class TwoPartyClient:
cdef public _Restorer _restorer cdef public _Restorer _restorer
cdef public _AsyncIoStream _stream cdef public _AsyncIoStream _stream
def __init__(self, socket, restorer=None): def __init__(self, socket, restorer=None, traversal_limit_in_words=None, nesting_limit=None):
if isinstance(socket, basestring): if isinstance(socket, basestring):
socket = self._connect(socket) socket = self._connect(socket)
cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._orig_stream = socket self._orig_stream = socket
self._stream = _FdAsyncIoStream(socket.fileno()) self._stream = _FdAsyncIoStream(socket.fileno())
self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.CLIENT) self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.CLIENT, opts)
if restorer is None: if restorer is None:
self.thisptr = new RpcSystem(makeRpcClient(deref(self._network.thisptr))) self.thisptr = new RpcSystem(makeRpcClient(deref(self._network.thisptr)))
self._restorer = None self._restorer = None
@@ -2340,11 +2350,13 @@ cdef class TwoPartyServer:
cdef capnp.TaskSet * _task_set cdef capnp.TaskSet * _task_set
cdef capnp.ErrorHandler _error_handler cdef capnp.ErrorHandler _error_handler
def __init__(self, socket, restorer=None, server_socket=None, bootstrap=None): def __init__(self, socket, restorer=None, server_socket=None, bootstrap=None,
traversal_limit_in_words=None, nesting_limit=None):
if not restorer and not bootstrap: if not restorer and not bootstrap:
raise KjException("You must provide either a bootstrap interface or a restorer (deperecated) to a server constructor.") raise KjException("You must provide either a bootstrap interface or a restorer (deperecated) to a server constructor.")
cdef _InterfaceSchema schema cdef _InterfaceSchema schema
cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._restorer = None self._restorer = None
self._bootstrap = None self._bootstrap = None
@@ -2355,7 +2367,7 @@ cdef class TwoPartyServer:
self._stream = _FdAsyncIoStream(socket.fileno()) self._stream = _FdAsyncIoStream(socket.fileno())
self._server_socket = server_socket self._server_socket = server_socket
self._port = 0 self._port = 0
self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER) self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER, opts)
if bootstrap: if bootstrap:
self._bootstrap = bootstrap self._bootstrap = bootstrap
@@ -3438,15 +3450,9 @@ cdef class _StreamFdMessageReader(_MessageReader):
:Parameters: - fd (`int`) - A file descriptor :Parameters: - fd (`int`) - A file descriptor
""" """
def __init__(self, file, traversal_limit_in_words = None, nesting_limit = None): def __init__(self, file, traversal_limit_in_words = None, nesting_limit = None):
cdef schema_cpp.ReaderOptions opts cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._parent = file self._parent = file
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
self.thisptr = new schema_cpp.StreamFdMessageReader(file.fileno(), opts) self.thisptr = new schema_cpp.StreamFdMessageReader(file.fileno(), opts)
def __dealloc__(self): def __dealloc__(self):
@@ -3469,15 +3475,9 @@ cdef class _PackedMessageReader(_MessageReader):
pass pass
cdef _init(self, schema_cpp.BufferedInputStream & stream, traversal_limit_in_words = None, nesting_limit = None, parent = None): cdef _init(self, schema_cpp.BufferedInputStream & stream, traversal_limit_in_words = None, nesting_limit = None, parent = None):
cdef schema_cpp.ReaderOptions opts cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._parent = parent self._parent = parent
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
self.thisptr = new schema_cpp.PackedMessageReader(stream, opts) self.thisptr = new schema_cpp.PackedMessageReader(stream, opts)
return self return self
@@ -3489,15 +3489,10 @@ cdef class _PackedMessageReaderBytes(_MessageReader):
cdef schema_cpp.ArrayInputStream * stream cdef schema_cpp.ArrayInputStream * stream
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 cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._parent = buf self._parent = buf
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
cdef const void *ptr cdef const void *ptr
cdef Py_ssize_t sz cdef Py_ssize_t sz
PyObject_AsReadBuffer(buf, &ptr, &sz) PyObject_AsReadBuffer(buf, &ptr, &sz)
@@ -3526,15 +3521,9 @@ cdef class _InputMessageReader(_MessageReader):
pass pass
cdef _init(self, schema_cpp.BufferedInputStream & stream, traversal_limit_in_words = None, nesting_limit = None, parent = None): cdef _init(self, schema_cpp.BufferedInputStream & stream, traversal_limit_in_words = None, nesting_limit = None, parent = None):
cdef schema_cpp.ReaderOptions opts cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._parent = parent self._parent = parent
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
self.thisptr = new schema_cpp.InputStreamMessageReader(stream, opts) self.thisptr = new schema_cpp.InputStreamMessageReader(stream, opts)
return self return self
@@ -3555,15 +3544,9 @@ cdef class _PackedFdMessageReader(_MessageReader):
:Parameters: - fd (`int`) - A file descriptor :Parameters: - fd (`int`) - A file descriptor
""" """
def __init__(self, file, traversal_limit_in_words = None, nesting_limit = None): def __init__(self, file, traversal_limit_in_words = None, nesting_limit = None):
cdef schema_cpp.ReaderOptions opts cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
self._parent = file self._parent = file
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
self.thisptr = new schema_cpp.PackedFdMessageReader(file.fileno(), opts) self.thisptr = new schema_cpp.PackedFdMessageReader(file.fileno(), opts)
def __dealloc__(self): def __dealloc__(self):
@@ -3757,14 +3740,9 @@ cdef class _BufferView:
cdef class _FlatArrayMessageReader(_MessageReader): cdef class _FlatArrayMessageReader(_MessageReader):
cdef object _object_to_pin cdef object _object_to_pin
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 cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
cdef _AlignedBuffer aligned cdef _AlignedBuffer aligned
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
sz = len(buf) sz = len(buf)
if sz % 8 != 0: if sz % 8 != 0:
raise ValueError("input length must be a multiple of eight bytes") raise ValueError("input length must be a multiple of eight bytes")
@@ -3796,11 +3774,7 @@ cdef class _SegmentArrayMessageReader(_MessageReader):
cdef schema_cpp.ConstWordArrayPtr* _seg_ptrs cdef schema_cpp.ConstWordArrayPtr* _seg_ptrs
def __init__(self, segments, traversal_limit_in_words = None, nesting_limit = None): def __init__(self, segments, traversal_limit_in_words = None, nesting_limit = None):
cdef schema_cpp.ReaderOptions opts cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit)
if traversal_limit_in_words is not None:
opts.traversalLimitInWords = traversal_limit_in_words
if nesting_limit is not None:
opts.nestingLimit = nesting_limit
# take a Python array of bytes and constructs a ConstWordArrayArrayPtr # take a Python array of bytes and constructs a ConstWordArrayArrayPtr
num_segments = len(segments) num_segments = len(segments)
cdef const void* ptr cdef const void* ptr

View File

@@ -43,6 +43,24 @@ def test_simple_rpc():
assert response.x == '125' assert response.x == '125'
def test_simple_rpc_with_options():
read, write = socket.socketpair(socket.AF_UNIX)
restorer = SimpleRestorer()
server = capnp.TwoPartyServer(write, restorer)
# This traversal limit is too low to receive the response in, so we expect
# an exception during the call.
client = capnp.TwoPartyClient(read, traversal_limit_in_words=1)
ref = test_capability_capnp.TestSturdyRefObjectId.new_message(tag='testInterface')
cap = client.restore(ref)
cap = cap.cast_as(test_capability_capnp.TestInterface)
remote = cap.foo(i=5)
with pytest.raises(capnp.KjException):
response = remote.wait()
def test_simple_rpc_restore_func(): def test_simple_rpc_restore_func():
read, write = socket.socketpair(socket.AF_UNIX) read, write = socket.socketpair(socket.AF_UNIX)