From 55dcff0dad323436627da91057b1d782f5719d08 Mon Sep 17 00:00:00 2001 From: Jason Paryani Date: Thu, 13 Feb 2014 22:08:26 -0800 Subject: [PATCH] Fix TwoPartyServer to handle more than one client at a time. Fixes #23 Also fix various memleaks in RpcServer/Client --- capnp/helpers/asyncHelper.h | 4 ++ capnp/helpers/helpers.pxd | 8 ++- capnp/helpers/non_circular.pxd | 2 + capnp/helpers/rpcHelper.h | 48 +++++++++++++++ capnp/includes/capnp_cpp.pxd | 5 +- capnp/lib/capnp.pyx | 106 +++++++++++++-------------------- 6 files changed, 105 insertions(+), 68 deletions(-) diff --git a/capnp/helpers/asyncHelper.h b/capnp/helpers/asyncHelper.h index 40711b8..9d3860a 100644 --- a/capnp/helpers/asyncHelper.h +++ b/capnp/helpers/asyncHelper.h @@ -27,3 +27,7 @@ public: private: PyObject * py_event_port; }; + +void waitNeverDone(kj::WaitScope & scope) { + kj::NEVER_DONE.wait(scope); +} diff --git a/capnp/helpers/helpers.pxd b/capnp/helpers/helpers.pxd index fa1ea21..5f1efa5 100644 --- a/capnp/helpers/helpers.pxd +++ b/capnp/helpers/helpers.pxd @@ -1,4 +1,4 @@ -from .capnp.includes.capnp_cpp cimport Maybe, DynamicStruct, Request, PyPromise, VoidPromise, PyPromiseArray, RemotePromise, DynamicCapability, InterfaceSchema, EnumSchema, StructSchema, DynamicValue, Capability, RpcSystem, MessageBuilder, MessageReader, TwoPartyVatNetwork, PyRestorer, AnyPointer, DynamicStruct_Builder +from .capnp.includes.capnp_cpp cimport Maybe, DynamicStruct, Request, PyPromise, VoidPromise, PyPromiseArray, RemotePromise, DynamicCapability, InterfaceSchema, EnumSchema, StructSchema, DynamicValue, Capability, RpcSystem, MessageBuilder, MessageReader, TwoPartyVatNetwork, PyRestorer, AnyPointer, DynamicStruct_Builder, WaitScope, AsyncIoContext, StringPtr, TaskSet from .capnp.includes.schema_cpp cimport ByteArray @@ -30,6 +30,10 @@ cdef extern from "../helpers/rpcHelper.h": Capability.Client restoreHelper(RpcSystem&, AnyPointer.Reader&) Capability.Client restoreHelper(RpcSystem&, AnyPointer.Builder&) RpcSystem makeRpcClientWithRestorer(TwoPartyVatNetwork&, PyRestorer&) + PyPromise connectServer(TaskSet &, PyRestorer &, AsyncIoContext *, StringPtr) cdef extern from "../helpers/serialize.h": - ByteArray messageToPackedBytes(MessageBuilder &) \ No newline at end of file + ByteArray messageToPackedBytes(MessageBuilder &) + +cdef extern from "../helpers/asyncHelper.h": + void waitNeverDone(WaitScope&) diff --git a/capnp/helpers/non_circular.pxd b/capnp/helpers/non_circular.pxd index 79e4043..49e5507 100644 --- a/capnp/helpers/non_circular.pxd +++ b/capnp/helpers/non_circular.pxd @@ -12,6 +12,8 @@ cdef extern from "../helpers/capabilityHelper.h": cdef extern from "../helpers/rpcHelper.h": cdef cppclass PyRestorer: PyRestorer(PyObject *) + cdef cppclass ErrorHandler: + pass cdef extern from "../helpers/asyncHelper.h": cdef cppclass PyEventPort: diff --git a/capnp/helpers/rpcHelper.h b/capnp/helpers/rpcHelper.h index b453067..b8ff922 100644 --- a/capnp/helpers/rpcHelper.h +++ b/capnp/helpers/rpcHelper.h @@ -65,3 +65,51 @@ capnp::RpcSystem makeRpcClientWithRestorer( return RpcSystem(network, kj::Maybe&>(restorer)); } + +struct ServerContext { + kj::Own stream; + capnp::TwoPartyVatNetwork network; + capnp::RpcSystem rpcSystem; + + ServerContext(kj::Own&& stream, capnp::SturdyRefRestorer& restorer) + : stream(kj::mv(stream)), + network(*this->stream, capnp::rpc::twoparty::Side::SERVER), + rpcSystem(makeRpcServer(network, restorer)) {} +}; + +class ErrorHandler : public kj::TaskSet::ErrorHandler { + void taskFailed(kj::Exception&& exception) override { + kj::throwFatalException(kj::mv(exception)); + } +}; + +void acceptLoop(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)); + + 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). + tasks.add(server->network.onDisconnect().attach(kj::mv(server))); + }))); +} + +kj::Promise connectServer(kj::TaskSet & tasks, PyRestorer & restorer, 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, + [&](kj::Own>&& portFulfiller, + kj::Own&& addr) { + auto listener = addr->listen(); + portFulfiller->fulfill(listener->getPort()); + acceptLoop(tasks, restorer, kj::mv(listener)); + }))); + + return portPromise.addBranch().then([&](uint port) { return PyLong_FromUnsignedLong(port); }); +} diff --git a/capnp/includes/capnp_cpp.pxd b/capnp/includes/capnp_cpp.pxd index c079070..6085ac5 100644 --- a/capnp/includes/capnp_cpp.pxd +++ b/capnp/includes/capnp_cpp.pxd @@ -5,7 +5,7 @@ cdef extern from "../helpers/checkCompiler.h": pass from schema_cpp cimport Node, Data, StructNode, EnumNode, InterfaceNode, MessageBuilder, MessageReader -from .capnp.helpers.non_circular cimport PythonInterfaceDynamicImpl, reraise_kj_exception, PyRefCounter, PyRestorer, PyEventPort +from .capnp.helpers.non_circular cimport PythonInterfaceDynamicImpl, reraise_kj_exception, PyRefCounter, PyRestorer, PyEventPort, ErrorHandler from .capnp.includes.types cimport * cdef extern from "capnp/common.h" namespace " ::capnp": @@ -108,6 +108,9 @@ cdef extern from "kj/async-io.h" namespace " ::kj": Own[AsyncIoProvider] provider WaitScope waitScope + cdef cppclass TaskSet: + TaskSet(ErrorHandler &) + AsyncIoContext setupAsyncIo() cdef extern from "capnp/schema.h" namespace " ::capnp": diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index f208ea0..d486008 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -1316,6 +1316,10 @@ cdef _EventLoop C_DEFAULT_EVENT_LOOP_GETTER(): # C_DEFAULT_EVENT_LOOP._remove() # C_DEFAULT_EVENT_LOOP = _EventLoop() +def wait_forever(): + cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() + helpers.waitNeverDone(deref(loop.thisptr).waitScope) + cdef class _CallContext: cdef CallContext * thisptr @@ -1771,17 +1775,8 @@ cdef class TwoPartyClient: self._restorer = _convert_restorer(restorer) self.thisptr = new RpcSystem(makeRpcClientWithRestorer(deref(self._network.thisptr), deref(self._restorer.thisptr))) - Py_INCREF(self._restorer) - Py_INCREF(self._orig_stream) - Py_INCREF(self._stream) - Py_INCREF(self._network) # TODO:MEMORY: attach this to onDrained, also figure out what's leaking - def __dealloc__(self): del self.thisptr - # Py_DECREF(self._restorer) - # Py_DECREF(self._orig_stream) - # Py_DECREF(self._stream) - # Py_DECREF(self._network) cpdef _connect(self, host_string): host, port = host_string.split(':') @@ -1831,89 +1826,70 @@ cdef class TwoPartyClient: return self.restore(ref) + cpdef on_disconnect(self) except +reraise_kj_exception: + return _VoidPromise()._init(deref(self._network.thisptr).onDisconnect()) + cdef class TwoPartyServer: cdef RpcSystem * thisptr cdef public _TwoPartyVatNetwork _network - cdef public object _orig_stream, _server_socket + cdef public object _orig_stream, _server_socket, _disconnect_promise cdef public _Restorer _restorer cdef public _FdAsyncIoStream _stream - cdef public int port + cdef object _port + cdef public object port_promise + cdef capnp.TaskSet * _task_set + cdef capnp.ErrorHandler _error_handler def __init__(self, socket, restorer, server_socket=None): + self._restorer = _convert_restorer(restorer) if isinstance(socket, basestring): self._connect(socket) else: self._orig_stream = socket self._stream = _FdAsyncIoStream(socket.fileno()) self._server_socket = server_socket - self.port = 0 - self._restorer = _convert_restorer(restorer) - self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER) - self.thisptr = new RpcSystem(makeRpcServer(deref(self._network.thisptr), deref(self._restorer.thisptr))) + self._port = 0 + self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER) + self.thisptr = new RpcSystem(makeRpcServer(deref(self._network.thisptr), deref(self._restorer.thisptr))) - Py_INCREF(self._orig_stream) - Py_INCREF(self._stream) - Py_INCREF(self._restorer) - Py_INCREF(self._network) # TODO:MEMORY: attach this to onDrained, also figure out what's leaking + Py_INCREF(self._orig_stream) + Py_INCREF(self._stream) + Py_INCREF(self._restorer) + Py_INCREF(self._network) + self._disconnect_promise = self.on_disconnect().then(self._decref) cpdef _connect(self, host_string): - if ':' in host_string: - address, port = host_string.split(':') - port = int(port) - else: - address = host_string - port = _random.randint(60000, 61000) + 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 address == '*': - address = '' - - s = _socket.socket(_socket.AF_INET, _socket.SOCK_STREAM) - - # Set TCP_NODELAY on socket to disable Nagle's algorithm. This is not - # neccessary, but it speeds things up. - s.setsockopt(_socket.IPPROTO_TCP, _socket.TCP_NODELAY, 1) - - s.bind((address, port)) - s.listen(1) # service only 1 client at a time - - (clientsocket, address) = s.accept() - - self._server_socket = s - self._orig_stream = clientsocket - self._stream = _FdAsyncIoStream(self._orig_stream.fileno()) - self.port = port + def _decref(self): + Py_DECREF(self._restorer) + Py_DECREF(self._orig_stream) + Py_DECREF(self._stream) + Py_DECREF(self._network) def __dealloc__(self): del self.thisptr - # Py_DECREF(self._restorer) - # Py_DECREF(self._orig_stream) - # Py_DECREF(self._stream) - # Py_DECREF(self._network) + del self._task_set cpdef on_disconnect(self) except +reraise_kj_exception: return _VoidPromise()._init(deref(self._network.thisptr).onDisconnect()) cpdef run_forever(self): - if self._server_socket is None: - raise ValueError("You must pass a `server_socket` parameter to __init__ or a string as the socket parameter to use this function") + if self.port_promise is None: + raise ValueError("You must pass a string as the socket parameter in __init__ to use this function") - while True: - try: - self.on_disconnect().wait() + wait_forever() - (clientsocket, address) = self._server_socket.accept() - self._orig_stream = clientsocket - self._stream = _FdAsyncIoStream(clientsocket.fileno()) - self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.SERVER) - - del self.thisptr - self.thisptr = new RpcSystem(makeRpcServer(deref(self._network.thisptr), deref(self._restorer.thisptr))) - - Py_INCREF(self._orig_stream) - Py_INCREF(self._stream) - Py_INCREF(self._network) # TODO:MEMORY: attach this to onDrained, also figure out what's leaking - except KeyboardInterrupt: - break + property port: + def __get__(self): + if self._port is None: + self._port = self.port_promise.wait() + return self._port + else: + return self._port # TODO: add restore functionality here?