From e00bff6ef85add1132119a2086176c4d6117b36e Mon Sep 17 00:00:00 2001 From: Jason Paryani Date: Mon, 13 Apr 2015 14:14:40 -0700 Subject: [PATCH] Fix problems with TwoPartyServer Fixes #61 It turns out I messed up the Server initialization code for the case where a string is passed in as the address. The tests only cover the cases where a raw socket is passed in. This will be rectified in a following commit. --- capnp/helpers/helpers.pxd | 3 +- capnp/helpers/rpcHelper.h | 57 ++++++++++++++++++++++++++++++----- capnp/lib/capnp.pyx | 14 +++++++-- examples/calculator_client.py | 2 +- examples/calculator_server.py | 7 +---- test/test_rpc_calculator.py | 4 +-- 6 files changed, 67 insertions(+), 20 deletions(-) diff --git a/capnp/helpers/helpers.pxd b/capnp/helpers/helpers.pxd index 28acaa1..15b7197 100644 --- a/capnp/helpers/helpers.pxd +++ b/capnp/helpers/helpers.pxd @@ -32,7 +32,8 @@ cdef extern from "capnp/helpers/rpcHelper.h": Capability.Client restoreHelper(RpcSystem&, AnyPointer.Builder&) Capability.Client bootstrapHelper(RpcSystem&) RpcSystem makeRpcClientWithRestorer(TwoPartyVatNetwork&, PyRestorer&) - PyPromise connectServer(TaskSet &, PyRestorer &, AsyncIoContext *, StringPtr) + PyPromise connectServerRestorer(TaskSet &, PyRestorer &, AsyncIoContext *, StringPtr) + PyPromise connectServer(TaskSet &, Capability.Client, AsyncIoContext *, StringPtr) cdef extern from "capnp/helpers/serialize.h": ByteArray messageToPackedBytes(MessageBuilder &, size_t wordCount) diff --git a/capnp/helpers/rpcHelper.h b/capnp/helpers/rpcHelper.h index 5ae53e3..cd25a92 100644 --- a/capnp/helpers/rpcHelper.h +++ b/capnp/helpers/rpcHelper.h @@ -89,12 +89,12 @@ capnp::RpcSystem makeRpcClientWithRestorer( return RpcSystem(network, restorer); } -struct ServerContext { +struct ServerContextRestorer { kj::Own stream; capnp::TwoPartyVatNetwork network; capnp::RpcSystem rpcSystem; - ServerContext(kj::Own&& stream, capnp::SturdyRefRestorer& restorer) + ServerContextRestorer(kj::Own&& stream, capnp::SturdyRefRestorer& restorer) : stream(kj::mv(stream)), network(*this->stream, capnp::rpc::twoparty::Side::SERVER), rpcSystem(makeRpcServer(network, restorer)) {} @@ -106,14 +106,14 @@ class ErrorHandler : public kj::TaskSet::ErrorHandler { } }; -void acceptLoop(kj::TaskSet & tasks, PyRestorer & restorer, kj::Own&& listener) { +void acceptLoopRestorer(kj::TaskSet & tasks, PyRestorer & restorer, kj::Own&& listener) { auto ptr = listener.get(); tasks.add(ptr->accept().then(kj::mvCapture(kj::mv(listener), [&](kj::Own&& listener, kj::Own&& connection) { - acceptLoop(tasks, restorer, kj::mv(listener)); + acceptLoopRestorer(tasks, restorer, kj::mv(listener)); - auto server = kj::heap(kj::mv(connection), restorer); + auto server = kj::heap(kj::mv(connection), restorer); // Arrange to destroy the server context when all references are gone, or when the // EzRpcServer is destroyed (which will destroy the TaskSet). @@ -121,7 +121,7 @@ void acceptLoop(kj::TaskSet & tasks, PyRestorer & restorer, kj::Own connectServer(kj::TaskSet & tasks, PyRestorer & restorer, kj::AsyncIoContext * context, kj::StringPtr bindAddress) { +kj::Promise connectServerRestorer(kj::TaskSet & tasks, PyRestorer & restorer, kj::AsyncIoContext * context, kj::StringPtr bindAddress) { auto paf = kj::newPromiseAndFulfiller(); auto portPromise = paf.promise.fork(); @@ -131,7 +131,50 @@ kj::Promise connectServer(kj::TaskSet & tasks, PyRestorer & restorer kj::Own&& addr) { auto listener = addr->listen(); portFulfiller->fulfill(listener->getPort()); - acceptLoop(tasks, restorer, kj::mv(listener)); + acceptLoopRestorer(tasks, restorer, kj::mv(listener)); + }))); + + return portPromise.addBranch().then([&](unsigned int port) { return PyLong_FromUnsignedLong(port); }); +} + + +struct ServerContext { + kj::Own stream; + capnp::TwoPartyVatNetwork network; + capnp::RpcSystem rpcSystem; + + ServerContext(kj::Own&& stream, capnp::Capability::Client client) + : stream(kj::mv(stream)), + network(*this->stream, capnp::rpc::twoparty::Side::SERVER), + rpcSystem(makeRpcServer(network, client)) {} +}; + +void acceptLoop(kj::TaskSet & tasks, capnp::Capability::Client client, kj::Own&& listener) { + auto ptr = listener.get(); + tasks.add(ptr->accept().then(kj::mvCapture(kj::mv(listener), + [&, client](kj::Own&& listener, + kj::Own&& connection) mutable { + acceptLoop(tasks, client, kj::mv(listener)); + + auto server = kj::heap(kj::mv(connection), client); + + // Arrange to destroy the server context when all references are gone, or when the + // EzRpcServer is destroyed (which will destroy the TaskSet). + tasks.add(server->network.onDisconnect().attach(kj::mv(server))); + }))); +} + +kj::Promise connectServer(kj::TaskSet & tasks, capnp::Capability::Client client, kj::AsyncIoContext * context, kj::StringPtr bindAddress) { + auto paf = kj::newPromiseAndFulfiller(); + auto portPromise = paf.promise.fork(); + + tasks.add(context->provider->getNetwork().parseAddress(bindAddress) + .then(kj::mvCapture(paf.fulfiller, + [&, client](kj::Own>&& portFulfiller, + kj::Own&& addr) mutable { + auto listener = addr->listen(); + portFulfiller->fulfill(listener->getPort()); + acceptLoop(tasks, client, kj::mv(listener)); }))); return portPromise.addBranch().then([&](unsigned int port) { return PyLong_FromUnsignedLong(port); }); diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index d16b4a7..3af300d 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -2279,7 +2279,7 @@ cdef class TwoPartyServer: self._bootstrap = None if isinstance(socket, basestring): - self._connect(socket) + self._connect(socket, restorer, bootstrap) else: self._orig_stream = socket self._stream = _FdAsyncIoStream(socket.fileno()) @@ -2302,11 +2302,19 @@ cdef class TwoPartyServer: Py_INCREF(self._network) self._disconnect_promise = self.on_disconnect().then(self._decref) - cpdef _connect(self, host_string): + cpdef _connect(self, host_string, restorer, bootstrap): + cdef _InterfaceSchema schema cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() cdef capnp.StringPtr temp_string = capnp.StringPtr(host_string, len(host_string)) self._task_set = new capnp.TaskSet(self._error_handler) - self.port_promise = Promise()._init(helpers.connectServer(deref(self._task_set), deref(self._restorer.thisptr), loop.thisptr, temp_string)) + if restorer: + self._restorer = _convert_restorer(restorer) + self.port_promise = Promise()._init(helpers.connectServerRestorer(deref(self._task_set), deref(self._restorer.thisptr), loop.thisptr, temp_string)) + else: + self._bootstrap = bootstrap + Py_INCREF(self._bootstrap) + schema = bootstrap.schema + self.port_promise = Promise()._init(helpers.connectServer(deref(self._task_set), helpers.server_to_client(schema.thisptr, bootstrap), loop.thisptr, temp_string)) def _decref(self): Py_DECREF(self._bootstrap) diff --git a/examples/calculator_client.py b/examples/calculator_client.py index 5dd1bd0..de246c3 100755 --- a/examples/calculator_client.py +++ b/examples/calculator_client.py @@ -39,7 +39,7 @@ def main(host): # takes a struct or AnyPointer as an argument), and then cast the returned # capability to it's proper type. This casting is due to capabilities not # having a reference to their schema - calculator = client.ez_restore('calculator').cast_as(calculator_capnp.Calculator) + calculator = client.bootstrap().cast_as(calculator_capnp.Calculator) '''Make a request that just evaluates the literal value 123. diff --git a/examples/calculator_server.py b/examples/calculator_server.py index ffd9867..f072f58 100755 --- a/examples/calculator_server.py +++ b/examples/calculator_server.py @@ -130,15 +130,10 @@ given address/port ADDRESS may be '*' to bind to all local addresses.\ return parser.parse_args() -def restore(ref): - assert ref.as_text() == 'calculator' - return CalculatorImpl() - - def main(): address = parse_args().address - server = capnp.TwoPartyServer(address, restore) + server = capnp.TwoPartyServer(address, bootstrap=CalculatorImpl()) server.run_forever() if __name__ == '__main__': diff --git a/test/test_rpc_calculator.py b/test/test_rpc_calculator.py index abc4b15..2eccf99 100644 --- a/test/test_rpc_calculator.py +++ b/test/test_rpc_calculator.py @@ -12,7 +12,7 @@ import calculator_server def test_calculator(): read, write = socket.socketpair(socket.AF_UNIX) - server = capnp.TwoPartyServer(write, calculator_server.restore) + server = capnp.TwoPartyServer(write, bootstrap=calculator_server.CalculatorImpl()) calculator_client.main(read) @@ -29,7 +29,7 @@ def test_calculator_gc(): evaluate_impl_orig = calculator_server.evaluate_impl calculator_server.evaluate_impl = new_evaluate_impl(evaluate_impl_orig) - server = capnp.TwoPartyServer(write, calculator_server.restore) + server = capnp.TwoPartyServer(write, bootstrap=calculator_server.CalculatorImpl()) calculator_client.main(read) calculator_server.evaluate_impl = evaluate_impl_orig