From 0382beb26a02292d5e6bd721f0a6608dfa20830d Mon Sep 17 00:00:00 2001 From: Alex Silverstein Date: Sun, 22 Jan 2017 01:57:01 +0000 Subject: [PATCH] Let TwoParty Clients and Servers take ReaderOptions --- capnp/includes/capnp_cpp.pxd | 6 +-- capnp/lib/capnp.pyx | 76 ++++++++++++------------------------ test/test_rpc.py | 18 +++++++++ 3 files changed, 46 insertions(+), 54 deletions(-) diff --git a/capnp/includes/capnp_cpp.pxd b/capnp/includes/capnp_cpp.pxd index 7816d20..e4f5675 100644 --- a/capnp/includes/capnp_cpp.pxd +++ b/capnp/includes/capnp_cpp.pxd @@ -4,7 +4,7 @@ cdef extern from "capnp/helpers/checkCompiler.h": 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.includes.types cimport * @@ -45,7 +45,7 @@ cdef extern from "kj/exception.h" namespace " ::kj": cdef extern from "kj/memory.h" namespace " ::kj": cdef cppclass Own[T]: 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 >"(PromiseFulfillerPair&) 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 cppclass TwoPartyVatNetwork: - TwoPartyVatNetwork(EventLoop &, AsyncIoStream& stream, Side) + TwoPartyVatNetwork(EventLoop &, AsyncIoStream& stream, Side, ReaderOptions) VoidPromise onDisconnect() VoidPromise onDrained() RpcSystem makeRpcServer(TwoPartyVatNetwork&, PyRestorer&) diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index 284bd7c..d73c15b 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -264,6 +264,14 @@ cdef public object get_exception_info(object exc_type, object exc_obj, object ex except: 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: _DynamicStructReader _DynamicStructBuilder @@ -2217,9 +2225,9 @@ cdef class _TwoPartyVatNetwork: cdef Own[C_TwoPartyVatNetwork] thisptr 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.thisptr = makeTwoPartyVatNetwork(deref(stream.thisptr), side) + self.thisptr = makeTwoPartyVatNetwork(deref(stream.thisptr), side, opts) return self cpdef on_disconnect(self) except +reraise_kj_exception: @@ -2244,13 +2252,15 @@ cdef class TwoPartyClient: cdef public _Restorer _restorer 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): socket = self._connect(socket) + cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit) + self._orig_stream = socket 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: self.thisptr = new RpcSystem(makeRpcClient(deref(self._network.thisptr))) self._restorer = None @@ -2340,11 +2350,13 @@ cdef class TwoPartyServer: cdef capnp.TaskSet * _task_set 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: raise KjException("You must provide either a bootstrap interface or a restorer (deperecated) to a server constructor.") cdef _InterfaceSchema schema + cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit) self._restorer = None self._bootstrap = None @@ -2355,7 +2367,7 @@ cdef class TwoPartyServer: self._stream = _FdAsyncIoStream(socket.fileno()) self._server_socket = server_socket self._port = 0 - self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER) + self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER, opts) if bootstrap: self._bootstrap = bootstrap @@ -3438,15 +3450,9 @@ cdef class _StreamFdMessageReader(_MessageReader): :Parameters: - fd (`int`) - A file descriptor """ 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 - - 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) def __dealloc__(self): @@ -3469,15 +3475,9 @@ cdef class _PackedMessageReader(_MessageReader): pass 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 - - 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) return self @@ -3489,15 +3489,10 @@ cdef class _PackedMessageReaderBytes(_MessageReader): cdef schema_cpp.ArrayInputStream * stream 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 - 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 Py_ssize_t sz PyObject_AsReadBuffer(buf, &ptr, &sz) @@ -3526,15 +3521,9 @@ cdef class _InputMessageReader(_MessageReader): pass 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 - - 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) return self @@ -3555,15 +3544,9 @@ cdef class _PackedFdMessageReader(_MessageReader): :Parameters: - fd (`int`) - A file descriptor """ 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 - - 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) def __dealloc__(self): @@ -3757,14 +3740,9 @@ cdef class _BufferView: cdef class _FlatArrayMessageReader(_MessageReader): cdef object _object_to_pin 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 - 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) if sz % 8 != 0: 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 def __init__(self, segments, traversal_limit_in_words = None, nesting_limit = None): - 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 + cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit) # take a Python array of bytes and constructs a ConstWordArrayArrayPtr num_segments = len(segments) cdef const void* ptr diff --git a/test/test_rpc.py b/test/test_rpc.py index d288d04..c51797e 100644 --- a/test/test_rpc.py +++ b/test/test_rpc.py @@ -43,6 +43,24 @@ def test_simple_rpc(): 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(): read, write = socket.socketpair(socket.AF_UNIX)