diff --git a/capnp/__init__.py b/capnp/__init__.py index b8389bb..0482b03 100644 --- a/capnp/__init__.py +++ b/capnp/__init__.py @@ -53,6 +53,7 @@ from .lib.capnp import ( _write_message_to_fd, _write_packed_message_to_fd, _Promise as Promise, + _AsyncIoStream as AsyncIoStream, _init_capnp_api, ) diff --git a/capnp/helpers/asyncHelper.h b/capnp/helpers/asyncHelper.h index 4d74e43..f939b3f 100644 --- a/capnp/helpers/asyncHelper.h +++ b/capnp/helpers/asyncHelper.h @@ -1,69 +1,13 @@ #pragma once #include "kj/async.h" -#include "Python.h" #include "capabilityHelper.h" -class PyEventPort: public kj::EventPort { -public: - PyEventPort(PyObject * _py_event_port): py_event_port(_py_event_port) { - // We don't need to incref/decref, since this C++ class will be owned by the Python wrapper class, and we'll make sure the python class doesn't refcount to 0 elsewhere. - // Py_INCREF(py_event_port); - } - virtual bool wait() { - GILAcquire gil; - PyObject_CallMethod(py_event_port, const_cast("wait"), NULL); - return true; // TODO: get the bool result from python - } - - virtual bool poll() { - GILAcquire gil; - PyObject_CallMethod(py_event_port, const_cast("poll"), NULL); - return true; // TODO: get the bool result from python - } - - virtual void setRunnable(bool runnable) { - GILAcquire gil; - PyObject * arg = Py_False; - if (runnable) - arg = Py_True; - PyObject_CallMethod(py_event_port, const_cast("set_runnable"), const_cast("o"), arg); - } - -private: - PyObject * py_event_port; -}; - void waitNeverDone(kj::WaitScope & scope) { - GILRelease gil; kj::NEVER_DONE.wait(scope); } -void pollWaitScope(kj::WaitScope & scope) { - GILRelease gil; - scope.poll(); -} - -kj::Timer * getTimer(kj::AsyncIoContext * context) { - return &context->lowLevelProvider->getTimer(); -} - -void waitVoidPromise(kj::Promise * promise, kj::WaitScope & scope) { - GILRelease gil; - promise->wait(scope); -} - -PyObject * waitPyPromise(kj::Promise * promise, kj::WaitScope & scope) { - GILRelease gil; - return promise->wait(scope); -} - -capnp::Response< ::capnp::DynamicStruct> * waitRemote(capnp::RemotePromise< ::capnp::DynamicStruct> * promise, kj::WaitScope & scope) { - GILRelease gil; +capnp::Response< ::capnp::DynamicStruct> * waitRemote(kj::Own> promise, + kj::WaitScope & scope) { return new capnp::Response< ::capnp::DynamicStruct>(promise->wait(scope)); } - -bool pollRemote(capnp::RemotePromise< ::capnp::DynamicStruct> * promise, kj::WaitScope & scope) { - GILRelease gil; - return promise->poll(scope); -} diff --git a/capnp/helpers/asyncIoHelper.h b/capnp/helpers/asyncIoHelper.h deleted file mode 100644 index c848ebb..0000000 --- a/capnp/helpers/asyncIoHelper.h +++ /dev/null @@ -1,53 +0,0 @@ -#pragma once - -#include "kj/async.h" -#include "kj/async-io.h" - -class AsyncIoStreamReadHelper { -public: - AsyncIoStreamReadHelper(kj::AsyncIoStream * _stream, kj::WaitScope * _scope, size_t bufsize) { - io_stream = _stream; - wait_scope = _scope; - ready = false; - buffer_read_size = 0; - buffer = new unsigned char[bufsize]; - promise = io_stream->read(buffer, 1, bufsize); - } - - ~AsyncIoStreamReadHelper() { - delete[] buffer; - } - - bool poll() { - bool result = promise.poll(*wait_scope); - if (result) { - ready = true; - buffer_read_size = promise.wait(*wait_scope); - } - return result; - } - - size_t read_size() { - if (!ready) { - return 0; - } - return buffer_read_size; - } - - void * read_buffer() { - if (!ready) { - return 0; - } - return buffer; - } - -private: - kj::AsyncIoStream * io_stream; - kj::WaitScope * wait_scope; - kj::Promise promise = nullptr; - - unsigned char *buffer; - size_t buffer_read_size; - - bool ready; -}; diff --git a/capnp/helpers/capabilityHelper.cpp b/capnp/helpers/capabilityHelper.cpp index 05350ac..92c0444 100644 --- a/capnp/helpers/capabilityHelper.cpp +++ b/capnp/helpers/capabilityHelper.cpp @@ -1,8 +1,9 @@ #include "capnp/helpers/capabilityHelper.h" #include "capnp/lib/capnp_api.h" -::kj::Promise convert_to_pypromise(capnp::RemotePromise & promise) { - return promise.then([](capnp::Response&& response) { return wrap_dynamic_struct_reader(response); } ); +::kj::Promise> convert_to_pypromise(kj::Own> promise) { + return promise->then([](capnp::Response&& response) { + return stealPyRef(wrap_dynamic_struct_reader(response)); } ); } void reraise_kj_exception() { @@ -57,85 +58,73 @@ void check_py_error() { } } -kj::Promise wrapPyFunc(PyObject * func, PyObject * arg) { +inline kj::Promise> maybeUnwrapPromise(PyObject * result) { + check_py_error(); + auto promise = extract_promise(result); + Py_DECREF(result); + return kj::mv(*promise); +} + +kj::Promise> wrapPyFunc(kj::Own func, kj::Own arg) { GILAcquire gil; - auto arg_promise = extract_promise(arg); - - if(arg_promise == NULL) { - PyObject * result = PyObject_CallFunctionObjArgs(func, arg, NULL); - Py_DECREF(arg); - - check_py_error(); - - auto promise = extract_promise(result); - if(promise != NULL) - return kj::mv(*promise); // TODO: delete promise, see incref of containing promise in capnp.pyx - auto remote_promise = extract_remote_promise(result); - if(remote_promise != NULL) - return convert_to_pypromise(*remote_promise); // TODO: delete promise, see incref of containing promise in capnp.pyx - return result; - } - else { - return arg_promise->then([&](PyObject * new_arg){ return wrapPyFunc(func, new_arg); });// TODO: delete arg_promise? - } + // Creates an owned reference, which will be destroyed in maybeUnwrapPromise + PyObject * result = PyObject_CallFunctionObjArgs(func->obj, arg->obj, NULL); + return maybeUnwrapPromise(result); } -kj::Promise wrapPyFuncNoArg(PyObject * func) { +kj::Promise> wrapPyFuncNoArg(kj::Own func) { GILAcquire gil; - PyObject * result = PyObject_CallFunctionObjArgs(func, NULL); - - check_py_error(); - - auto promise = extract_promise(result); - if(promise != NULL) - return kj::mv(*promise); - auto remote_promise = extract_remote_promise(result); - if(remote_promise != NULL) - return convert_to_pypromise(*remote_promise); // TODO: delete promise, see incref of containing promise in capnp.pyx - return result; + // Creates an owned reference, which will be destroyed in maybeUnwrapPromise + PyObject * result = PyObject_CallFunctionObjArgs(func->obj, NULL); + return maybeUnwrapPromise(result); } -kj::Promise wrapRemoteCall(PyObject * func, capnp::Response & arg) { +kj::Promise> wrapRemoteCall(kj::Own func, capnp::Response & arg) { GILAcquire gil; - PyObject * ret = wrap_remote_call(func, arg); - - check_py_error(); - - auto promise = extract_promise(ret); - if(promise != NULL) - return kj::mv(*promise); - auto remote_promise = extract_remote_promise(ret); - if(remote_promise != NULL) - return convert_to_pypromise(*remote_promise); // TODO: delete promise, see incref of containing promise in capnp.pyx - return ret; + // Creates an owned reference, which will be destroyed in maybeUnwrapPromise + PyObject * ret = wrap_remote_call(func->obj, arg); + return maybeUnwrapPromise(ret); } -::kj::Promise then(kj::Promise & promise, PyObject * func, PyObject * error_func) { - if(error_func == Py_None) - return promise.then([func](PyObject * arg) { return wrapPyFunc(func, arg); } ); +::kj::Promise> then(kj::Own>> promise, + kj::Own func, kj::Own error_func) { + if(error_func->obj == Py_None) + return promise->then(kj::mvCapture(func, [](auto func, kj::Own arg) { + return wrapPyFunc(kj::mv(func), kj::mv(arg)); } )); else - return promise.then([func](PyObject * arg) { return wrapPyFunc(func, arg); } - , [error_func](kj::Exception arg) { return wrapPyFunc(error_func, wrap_kj_exception(arg)); } ); + return promise->then + (kj::mvCapture(func, [](auto func, kj::Own arg) { + return wrapPyFunc(kj::mv(func), kj::mv(arg)); }), + kj::mvCapture(error_func, [](auto error_func, kj::Exception arg) { + return wrapPyFunc(kj::mv(error_func), stealPyRef(wrap_kj_exception(arg))); } )); } -::kj::Promise then(::capnp::RemotePromise< ::capnp::DynamicStruct> & promise, PyObject * func, PyObject * error_func) { - if(error_func == Py_None) - return promise.then([func](capnp::Response&& arg) { return wrapRemoteCall(func, arg); } ); +::kj::Promise> then(kj::Own<::capnp::RemotePromise<::capnp::DynamicStruct>> promise, + kj::Own func, kj::Own error_func) { + if(error_func->obj == Py_None) + return promise->then(kj::mvCapture(func, [](auto func, capnp::Response&& arg) { + return wrapRemoteCall(kj::mv(func), arg); } )); else - return promise.then([func](capnp::Response&& arg) { return wrapRemoteCall(func, arg); } - , [error_func](kj::Exception arg) { return wrapPyFunc(error_func, wrap_kj_exception(arg)); } ); + return promise->then + (kj::mvCapture(func, [](auto func, capnp::Response&& arg) { + return wrapRemoteCall(kj::mv(func), arg); }), + kj::mvCapture(error_func, [](auto error_func, kj::Exception arg) { + return wrapPyFunc(kj::mv(error_func), stealPyRef(wrap_kj_exception(arg))); } )); } -::kj::Promise then(kj::Promise & promise, PyObject * func, PyObject * error_func) { - if(error_func == Py_None) - return promise.then([func]() { return wrapPyFuncNoArg(func); } ); +::kj::Promise> then(kj::Own> promise, + kj::Own func, kj::Own error_func) { + if(error_func->obj == Py_None) + return promise->then(kj::mvCapture(func, [](auto func) { return wrapPyFuncNoArg(kj::mv(func)); } )); else - return promise.then([func]() { return wrapPyFuncNoArg(func); } - , [error_func](kj::Exception arg) { return wrapPyFunc(error_func, wrap_kj_exception(arg)); } ); + return promise->then(kj::mvCapture(func, [](auto func) { return wrapPyFuncNoArg(kj::mv(func)); }), + kj::mvCapture(error_func, [](auto error_func, kj::Exception arg) { + return wrapPyFunc(kj::mv(error_func), stealPyRef(wrap_kj_exception(arg))); } )); } -::kj::Promise then(kj::Promise > && promise) { - return promise.then([](kj::Array&& arg) { return convert_array_pyobject(arg); } ); +::kj::Promise> then(kj::Promise> > && promise) { + return promise.then([](kj::Array>&& arg) { + return stealPyRef(convert_array_pyobject(arg)); } ); } kj::Promise PythonInterfaceDynamicImpl::call(capnp::InterfaceSchema::Method method, @@ -154,6 +143,66 @@ kj::Promise PythonInterfaceDynamicImpl::call(capnp::InterfaceSchema::Metho return ret; }; + +class ReadPromiseAdapter { +public: + ReadPromiseAdapter(kj::PromiseFulfiller& fulfiller, PyObject* protocol, + void* buffer, size_t minBytes, size_t maxBytes) + : protocol(protocol) { + _asyncio_stream_read_start(protocol, buffer, minBytes, maxBytes, fulfiller); + } + + ~ReadPromiseAdapter() { + _asyncio_stream_read_stop(protocol); + } + +private: + PyObject* protocol; +}; + + +class WritePromiseAdapter { +public: + WritePromiseAdapter(kj::PromiseFulfiller& fulfiller, PyObject* protocol, + kj::ArrayPtr> pieces) + : protocol(protocol) { + _asyncio_stream_write_start(protocol, pieces, fulfiller); + } + + ~WritePromiseAdapter() { + _asyncio_stream_write_stop(protocol); + } + +private: + PyObject* protocol; + +}; + +PyAsyncIoStream::~PyAsyncIoStream() { + _asyncio_stream_close(protocol->obj); +} + +kj::Promise PyAsyncIoStream::tryRead(void* buffer, size_t minBytes, size_t maxBytes) { + return kj::newAdaptedPromise(protocol->obj, buffer, minBytes, maxBytes); +} + +kj::Promise PyAsyncIoStream::write(const void* buffer, size_t size) { + KJ_UNIMPLEMENTED("No use-case AsyncIoStream::write was found yet."); +} + +kj::Promise PyAsyncIoStream::write(kj::ArrayPtr> pieces) { + return kj::newAdaptedPromise(protocol->obj, pieces); +} + +kj::Promise PyAsyncIoStream::whenWriteDisconnected() { + // TODO: Possibly connect this to protocol.connection_lost? + return kj::NEVER_DONE; +} + +void PyAsyncIoStream::shutdownWrite() { + _asyncio_stream_shutdown_write(protocol->obj); +} + void init_capnp_api() { import_capnp__lib__capnp(); } diff --git a/capnp/helpers/capabilityHelper.h b/capnp/helpers/capabilityHelper.h index d45c511..d41bad9 100644 --- a/capnp/helpers/capabilityHelper.h +++ b/capnp/helpers/capabilityHelper.h @@ -1,6 +1,7 @@ #pragma once #include "capnp/dynamic.h" +#include #include #include "Python.h" @@ -26,57 +27,6 @@ public: PyThreadState *_save; // The macros above read/write from this variable }; -::kj::Promise convert_to_pypromise(capnp::RemotePromise & promise); - -inline ::kj::Promise convert_to_pypromise(kj::Promise & promise) { - return promise.then([]() { - GILAcquire gil; - Py_INCREF( Py_None ); - return Py_None; - }); -} - -template -::kj::Promise convert_to_voidpromise(kj::Promise & promise) { - return promise.then([](T) { } ); -} - -void reraise_kj_exception(); - -void check_py_error(); - -kj::Promise wrapPyFunc(PyObject * func, PyObject * arg); - -kj::Promise wrapPyFuncNoArg(PyObject * func); - -kj::Promise wrapRemoteCall(PyObject * func, capnp::Response & arg); - -::kj::Promise then(kj::Promise & promise, PyObject * func, PyObject * error_func); -::kj::Promise then(::capnp::RemotePromise< ::capnp::DynamicStruct> & promise, PyObject * func, PyObject * error_func); - -::kj::Promise then(kj::Promise & promise, PyObject * func, PyObject * error_func); - -::kj::Promise then(kj::Promise > && promise); - -class PythonInterfaceDynamicImpl final: public capnp::DynamicCapability::Server { -public: - PyObject * py_server; - - PythonInterfaceDynamicImpl(capnp::InterfaceSchema & schema, PyObject * _py_server) - : capnp::DynamicCapability::Server(schema), py_server(_py_server) { - GILAcquire gil; - Py_INCREF(_py_server); - } - - ~PythonInterfaceDynamicImpl() { - GILAcquire gil; - Py_DECREF(py_server); - } - - kj::Promise call(capnp::InterfaceSchema::Method method, - capnp::CallContext< capnp::DynamicStruct, capnp::DynamicStruct> context); -}; - class PyRefCounter { public: PyObject * obj; @@ -97,6 +47,63 @@ public: } }; +inline kj::Own stealPyRef(PyObject* o) { + auto ret = kj::heap(o); + Py_DECREF(o); + return ret; +} + +::kj::Promise> convert_to_pypromise(kj::Own> promise); + +inline ::kj::Promise> convert_to_pypromise(kj::Own> promise) { + return promise->then([]() { + GILAcquire gil; + return kj::heap(Py_None); + }); +} + +template +::kj::Promise convert_to_voidpromise(kj::Own> promise) { + return promise->then([](T) { } ); +} + +void reraise_kj_exception(); + +void check_py_error(); + +inline kj::Promise> wrapSizePromise(kj::Promise promise) { + return promise.then([](size_t response) { return stealPyRef(PyLong_FromSize_t(response)); } ); +} + +::kj::Promise> then(kj::Own>> promise, + kj::Own func, kj::Own error_func); +::kj::Promise> then(kj::Own<::capnp::RemotePromise< ::capnp::DynamicStruct>> promise, + kj::Own func, kj::Own error_func); + +::kj::Promise> then(kj::Own> promise, + kj::Ownfunc, kj::Own error_func); + +::kj::Promise> then(kj::Promise> > && promise); + +class PythonInterfaceDynamicImpl final: public capnp::DynamicCapability::Server { +public: + PyObject * py_server; + + PythonInterfaceDynamicImpl(capnp::InterfaceSchema & schema, PyObject * _py_server) + : capnp::DynamicCapability::Server(schema), py_server(_py_server) { + GILAcquire gil; + Py_INCREF(_py_server); + } + + ~PythonInterfaceDynamicImpl() { + GILAcquire gil; + Py_DECREF(py_server); + } + + kj::Promise call(capnp::InterfaceSchema::Method method, + capnp::CallContext< capnp::DynamicStruct, capnp::DynamicStruct> context); +}; + inline capnp::DynamicCapability::Client new_client(capnp::InterfaceSchema & schema, PyObject * server) { return capnp::DynamicCapability::Client(kj::heap(schema, server)); } @@ -108,4 +115,30 @@ inline capnp::Capability::Client server_to_client(capnp::InterfaceSchema & schem return kj::heap(schema, server); } +class PyAsyncIoStream: public kj::AsyncIoStream { +public: + kj::Own protocol; + + PyAsyncIoStream(kj::Own protocol) : protocol(kj::mv(protocol)) {} + ~PyAsyncIoStream(); + + kj::Promise tryRead(void* buffer, size_t minBytes, size_t maxBytes); + + kj::Promise write(const void* buffer, size_t size); + + kj::Promise write(kj::ArrayPtr> pieces); + + kj::Promise whenWriteDisconnected(); + + void shutdownWrite(); +}; + +template +inline void rejectDisconnected(kj::PromiseFulfiller& fulfiller, kj::StringPtr message) { + fulfiller.reject(KJ_EXCEPTION(DISCONNECTED, message)); +} +inline void rejectVoidDisconnected(kj::PromiseFulfiller& fulfiller, kj::StringPtr message) { + fulfiller.reject(KJ_EXCEPTION(DISCONNECTED, message)); +} + void init_capnp_api(); diff --git a/capnp/helpers/helpers.pxd b/capnp/helpers/helpers.pxd index fa55cd7..45e3f4c 100644 --- a/capnp/helpers/helpers.pxd +++ b/capnp/helpers/helpers.pxd @@ -1,8 +1,9 @@ from capnp.includes.capnp_cpp cimport ( - Maybe, ReaderOptions, DynamicStruct, Request, Response, PyPromise, VoidPromise, PyPromiseArray, + Maybe, ReaderOptions, DynamicStruct, Request, Response, Promise, PyPromise, VoidPromise, PyPromiseArray, RemotePromise, DynamicCapability, InterfaceSchema, EnumSchema, StructSchema, DynamicValue, Capability, RpcSystem, MessageBuilder, MessageReader, TwoPartyVatNetwork, AnyPointer, - DynamicStruct_Builder, WaitScope, AsyncIoContext, StringPtr, TaskSet, Timer, AsyncIoStreamReadHelper, + DynamicStruct_Builder, WaitScope, AsyncIoContext, StringPtr, TaskSet, Timer, + LowLevelAsyncIoProvider, AsyncIoProvider, Own, PyRefCounter ) from capnp.includes.schema_cpp cimport ByteArray @@ -20,31 +21,27 @@ cdef extern from "capnp/helpers/fixMaybe.h": cdef extern from "capnp/helpers/capabilityHelper.h": # PyPromise evalLater(EventLoop &, PyObject * func) # PyPromise there(EventLoop & loop, PyPromise & promise, PyObject * func, PyObject * error_func) - PyPromise then(PyPromise & promise, PyObject * func, PyObject * error_func) - PyPromise then(RemotePromise & promise, PyObject * func, PyObject * error_func) - PyPromise then(VoidPromise & promise, PyObject * func, PyObject * error_func) + PyPromise then(Own[PyPromise] promise, Own[PyRefCounter] func, Own[PyRefCounter] error_func) + PyPromise then(Own[RemotePromise] promise, Own[PyRefCounter] func, Own[PyRefCounter] error_func) + PyPromise then(Own[VoidPromise] promise, Own[PyRefCounter] func, Own[PyRefCounter] error_func) PyPromise then(PyPromiseArray & promise) DynamicCapability.Client new_client(InterfaceSchema&, PyObject *) DynamicValue.Reader new_server(InterfaceSchema&, PyObject *) Capability.Client server_to_client(InterfaceSchema&, PyObject *) - PyPromise convert_to_pypromise(RemotePromise&) - PyPromise convert_to_pypromise(VoidPromise&) - VoidPromise convert_to_voidpromise(PyPromise&) + PyPromise convert_to_pypromise(Own[RemotePromise]) + PyPromise convert_to_pypromise(Own[VoidPromise]) + VoidPromise convert_to_voidpromise(Own[PyPromise]) + PyPromise wrapSizePromise(Promise[size_t]) void init_capnp_api() cdef extern from "capnp/helpers/rpcHelper.h": Capability.Client bootstrapHelper(RpcSystem&) Capability.Client bootstrapHelperServer(RpcSystem&) - PyPromise connectServer(TaskSet &, Capability.Client, AsyncIoContext *, StringPtr, ReaderOptions &) + PyPromise connectServer(TaskSet &, Capability.Client, AsyncIoProvider *, StringPtr, ReaderOptions &) cdef extern from "capnp/helpers/serialize.h": ByteArray messageToPackedBytes(MessageBuilder &, size_t wordCount) cdef extern from "capnp/helpers/asyncHelper.h": - void waitNeverDone(WaitScope&) - void pollWaitScope(WaitScope&) - Response * waitRemote(RemotePromise *, WaitScope&) - bool pollRemote(RemotePromise *, WaitScope&) - PyObject * waitPyPromise(PyPromise *, WaitScope&) - void waitVoidPromise(VoidPromise *, WaitScope&) - Timer * getTimer(AsyncIoContext *) except +reraise_kj_exception + void waitNeverDone(WaitScope&) except +reraise_kj_exception nogil + Response * waitRemote(Own[RemotePromise], WaitScope&) except +reraise_kj_exception nogil diff --git a/capnp/helpers/non_circular.pxd b/capnp/helpers/non_circular.pxd index da3793b..a6246b3 100644 --- a/capnp/helpers/non_circular.pxd +++ b/capnp/helpers/non_circular.pxd @@ -9,11 +9,8 @@ cdef extern from "capnp/helpers/capabilityHelper.h": void reraise_kj_exception() cdef cppclass PyRefCounter: PyRefCounter(PyObject *) + PyObject * obj cdef extern from "capnp/helpers/rpcHelper.h": cdef cppclass ErrorHandler: pass - -cdef extern from "capnp/helpers/asyncHelper.h": - cdef cppclass PyEventPort: - PyEventPort(PyObject *) diff --git a/capnp/helpers/rpcHelper.h b/capnp/helpers/rpcHelper.h index 6d903e8..a0c68e4 100644 --- a/capnp/helpers/rpcHelper.h +++ b/capnp/helpers/rpcHelper.h @@ -52,11 +52,11 @@ void acceptLoop(kj::TaskSet & tasks, capnp::Capability::Client client, kj::Own connectServer(kj::TaskSet & tasks, capnp::Capability::Client client, kj::AsyncIoContext * context, kj::StringPtr bindAddress, capnp::ReaderOptions & opts) { +kj::Promise> connectServer(kj::TaskSet & tasks, capnp::Capability::Client client, kj::AsyncIoProvider * provider, kj::StringPtr bindAddress, capnp::ReaderOptions & opts) { auto paf = kj::newPromiseAndFulfiller(); auto portPromise = paf.promise.fork(); - tasks.add(context->provider->getNetwork().parseAddress(bindAddress) + tasks.add(provider->getNetwork().parseAddress(bindAddress) .then(kj::mvCapture(paf.fulfiller, [&, client, opts](kj::Own>&& portFulfiller, kj::Own&& addr) mutable { @@ -65,5 +65,6 @@ kj::Promise connectServer(kj::TaskSet & tasks, capnp::Capability::Cl acceptLoop(tasks, client, kj::mv(listener), opts); }))); - return portPromise.addBranch().then([&](unsigned int port) { return PyLong_FromUnsignedLong(port); }); + return portPromise.addBranch().then([&](unsigned int port) { + return stealPyRef(PyLong_FromUnsignedLong(port)); }); } diff --git a/capnp/includes/capnp_cpp.pxd b/capnp/includes/capnp_cpp.pxd index 17a9993..346533c 100644 --- a/capnp/includes/capnp_cpp.pxd +++ b/capnp/includes/capnp_cpp.pxd @@ -5,7 +5,7 @@ cdef extern from "capnp/helpers/checkCompiler.h": from libcpp cimport bool from capnp.helpers.non_circular cimport ( - PythonInterfaceDynamicImpl, reraise_kj_exception, PyRefCounter, PyEventPort, ErrorHandler, + PythonInterfaceDynamicImpl, reraise_kj_exception, PyRefCounter, ErrorHandler, ) from capnp.includes.schema_cpp cimport ( Node, Data, StructNode, EnumNode, InterfaceNode, MessageBuilder, MessageReader, ReaderOptions, @@ -48,21 +48,21 @@ cdef extern from "kj/exception.h" namespace " ::kj": cdef extern from "kj/memory.h" namespace " ::kj": cdef cppclass Own[T] nogil: + Own() T& operator*() T* get() + Own[T] heap[T](...) 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 *) cdef extern from "kj/async.h" namespace " ::kj": cdef cppclass Promise[T] nogil: - Promise() Promise(Promise) Promise(T) - T wait(WaitScope) - bool poll(WaitScope) + T wait(WaitScope) except +reraise_kj_exception + bool poll(WaitScope) except +reraise_kj_exception # ForkedPromise fork() # Promise exclusiveJoin(Promise&& other) # Promise[T] eagerlyEvaluate() @@ -73,7 +73,15 @@ cdef extern from "kj/async.h" namespace " ::kj": Promise[T] attach(Own[PyRefCounter] &, Own[PyRefCounter] &, Own[PyRefCounter] &) Promise[T] attach(Own[PyRefCounter] &, Own[PyRefCounter] &, Own[PyRefCounter] &, Own[PyRefCounter] &) -ctypedef Promise[PyObject *] PyPromise + cdef cppclass Canceler nogil: + Canceler() + Promise[T] wrap[T](Promise[T] promise) + void cancel(StringPtr cancelReason) + void cancel(Exception& exception) + void release() + bool isEmpty() + +ctypedef Promise[Own[PyRefCounter]] PyPromise ctypedef Promise[void] VoidPromise cdef extern from "kj/string-tree.h" namespace " ::kj": @@ -82,10 +90,11 @@ cdef extern from "kj/string-tree.h" namespace " ::kj": cdef extern from "kj/common.h" namespace " ::kj": cdef cppclass Maybe[T] nogil: - pass + T& orDefault(T&) cdef cppclass ArrayPtr[T] nogil: ArrayPtr() ArrayPtr(T *, size_t size) + T* begin() size_t size() T& operator[](size_t index) @@ -101,9 +110,9 @@ cdef extern from "kj/array.h" namespace " ::kj": T& add(T&) Array[T] finish() - ArrayBuilder[PyPromise] heapArrayBuilderPyPromise"::kj::heapArrayBuilder< ::kj::Promise >"(size_t) nogil + ArrayBuilder[PyPromise] heapArrayBuilderPyPromise"::kj::heapArrayBuilder< ::kj::Promise> >"(size_t) nogil - ctypedef Array[PyObject *] PyArray' ::kj::Array' + ctypedef Array[Own[PyRefCounter]] PyArray' ::kj::Array>' ctypedef Promise[PyArray] PyPromiseArray @@ -117,12 +126,23 @@ cdef extern from "kj/time.h" namespace " ::kj": Duration MINUTES Duration HOURS Duration DAYS - # cdef cppclass TimePoint: - # TimePoint(Duration) + cdef cppclass TimePoint: + TimePoint(Duration) + cdef cppclass MonotonicClock nogil: + MonotonicClock(MonotonicClock&) + TimePoint now() + MonotonicClock systemPreciseMonotonicClock() + +cdef extern from "kj/timer.h" namespace " ::kj": cdef cppclass Timer nogil: # int64_t now() # VoidPromise atTime(TimePoint time) VoidPromise afterDelay(Duration delay) + cdef cppclass TimerImpl(Timer) nogil: + TimerImpl(TimePoint startTime) + Maybe[TimePoint] nextEvent() + Maybe[uint64_t] timeoutToNextEvent(TimePoint start, Duration unit, uint64_t max) + void advanceTo(TimePoint newTime) cdef inline Duration Nanoseconds(int64_t nanos): return NANOSECONDS * nanos @@ -132,7 +152,7 @@ cdef extern from "kj/async-io.h" namespace " ::kj": Promise[size_t] read(void*, size_t, size_t) Promise[void] write(const void*, size_t) - cdef cppclass LowLevelAsyncIoProvider nogil: + cdef cppclass LowLevelAsyncIoProvider: # Own[AsyncInputStream] wrapInputFd(int) # Own[AsyncOutputStream] wrapOutputFd(int) Own[AsyncIoStream] wrapSocketFd(int) @@ -141,9 +161,6 @@ cdef extern from "kj/async-io.h" namespace " ::kj": cdef cppclass AsyncIoProvider nogil: TwoWayPipe newTwoWayPipe() - cdef cppclass WaitScope nogil: - pass - cdef cppclass AsyncIoContext nogil: AsyncIoContext(AsyncIoContext&) Own[LowLevelAsyncIoProvider] lowLevelProvider @@ -157,6 +174,7 @@ cdef extern from "kj/async-io.h" namespace " ::kj": Own[AsyncIoStream] ends[2] AsyncIoContext setupAsyncIo() nogil + Own[AsyncIoProvider] newAsyncIoProvider(LowLevelAsyncIoProvider& lowLevel); cdef extern from "capnp/schema.capnp.h" namespace " ::capnp": enum TypeWhich" ::capnp::schema::Type::Which": @@ -518,20 +536,31 @@ cdef extern from "capnp/capability.h" namespace " ::capnp": void allowCancellation() except +reraise_kj_exception cdef extern from "kj/async.h" namespace " ::kj": + cdef cppclass EventPort: + bool wait() except* with gil + bool poll() except* with gil + void setRunnable(bool runnable) except* with gil cdef cppclass EventLoop nogil: EventLoop() - EventLoop(PyEventPort &) - cdef cppclass PromiseFulfiller nogil: + EventLoop(EventPort &) + void run() + cdef cppclass WaitScope nogil: + WaitScope(EventLoop &) + void poll() + cdef cppclass PromiseFulfiller[T] nogil: + void fulfill(T&& value) + void reject(Exception&& exception) + cdef cppclass VoidPromiseFulfiller"::kj::PromiseFulfiller" nogil: void fulfill() + void reject(Exception&& exception) cdef cppclass PromiseFulfillerPair" ::kj::PromiseFulfillerPair" nogil: VoidPromise promise - Own[PromiseFulfiller] fulfiller + Own[VoidPromiseFulfiller] fulfiller PromiseFulfillerPair newPromiseAndFulfiller" ::kj::newPromiseAndFulfiller"() nogil PyPromiseArray joinPromises(Array[PyPromise]) nogil -cdef extern from "capnp/helpers/asyncIoHelper.h": - cdef cppclass AsyncIoStreamReadHelper nogil: - AsyncIoStreamReadHelper(AsyncIoStream *, WaitScope *, size_t) - bool poll() - size_t read_size() - void* read_buffer() +cdef extern from "capnp/helpers/capabilityHelper.h": + cdef cppclass PyAsyncIoStream(AsyncIoStream): + PyAsyncIoStream(PyObject* thisptr) + void rejectDisconnected[T](PromiseFulfiller[T]& fulfiller, StringPtr message) + void rejectVoidDisconnected(VoidPromiseFulfiller& fulfiller, StringPtr message) diff --git a/capnp/lib/capnp.pxd b/capnp/lib/capnp.pxd index 3dabdbd..21a3e41 100644 --- a/capnp/lib/capnp.pxd +++ b/capnp/lib/capnp.pxd @@ -12,7 +12,7 @@ from capnp.includes.capnp_cpp cimport ( CallContext, RpcSystem, makeRpcServerBootstrap, makeRpcClient, Capability as C_Capability, TwoPartyVatNetwork as C_TwoPartyVatNetwork, Side, AsyncIoStream, Own, makeTwoPartyVatNetwork, PromiseFulfillerPair as C_PromiseFulfillerPair, copyPromiseFulfillerPair, newPromiseAndFulfiller, - PyArray, DynamicStruct_Builder, TwoWayPipe, + PyArray, DynamicStruct_Builder, TwoWayPipe, PyRefCounter, PyAsyncIoStream ) from capnp.includes.schema_cpp cimport Node as C_Node, EnumNode as C_EnumNode from capnp.includes.types cimport * @@ -159,12 +159,9 @@ cdef _setDynamicFieldWithField(DynamicStruct_Builder thisptr, _StructSchemaField cdef _setDynamicFieldStatic(DynamicStruct_Builder thisptr, field, value, parent) cdef api object wrap_dynamic_struct_reader(Response & r) with gil -cdef api PyObject * wrap_remote_call(PyObject * func, Response & r) except * with gil cdef api Promise[void] * call_server_method( PyObject * _server, char * _method_name, CallContext & _context) except * with gil cdef api convert_array_pyobject(PyArray & arr) with gil -cdef api Promise[PyObject*] * extract_promise(object obj) with gil -cdef api RemotePromise * extract_remote_promise(object obj) with gil cdef api object wrap_kj_exception(capnp.Exception & exception) with gil cdef api object wrap_kj_exception_for_reraise(capnp.Exception & exception) with gil cdef api object get_exception_info(object exc_type, object exc_obj, object exc_tb) with gil diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index cebb0ed..c8c758d 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -9,15 +9,17 @@ cimport cython # noqa: E402 -from capnp.helpers.helpers cimport AsyncIoStreamReadHelper, init_capnp_api -from capnp.includes.capnp_cpp cimport AsyncIoStream, WaitScope, PyPromise, VoidPromise +from capnp.helpers.helpers cimport init_capnp_api +from capnp.includes.capnp_cpp cimport AsyncIoStream, WaitScope, PyPromise, VoidPromise, EventPort, EventLoop, WaitScope, LowLevelAsyncIoProvider, AsyncIoProvider, newAsyncIoProvider, MonotonicClock, Timer, TimerImpl, systemPreciseMonotonicClock, MILLISECONDS, Canceler, PyAsyncIoStream, PromiseFulfiller, VoidPromiseFulfiller -from cpython cimport array, Py_buffer, PyObject_CheckBuffer +from cpython cimport array, Py_buffer, PyObject_CheckBuffer, memoryview, buffer from cpython.buffer cimport PyBUF_SIMPLE, PyBUF_WRITABLE from cpython.exc cimport PyErr_Clear from cython.operator cimport dereference as deref from libc.stdlib cimport malloc, free from libc.string cimport memcpy +from libcpp.utility cimport move + import array import asyncio @@ -32,6 +34,7 @@ import sys as _sys import threading as _threading import traceback as _traceback import warnings as _warnings +import weakref as _weakref from types import ModuleType as _ModuleType from operator import attrgetter as _attrgetter @@ -58,14 +61,9 @@ cdef api object wrap_dynamic_struct_reader(Response & r) with gil: return _Response()._init_childptr(new Response(moveResponse(r)), None) -cdef api PyObject * wrap_remote_call(PyObject * func, Response & r) except * with gil: +cdef api object wrap_remote_call(object func, Response & r): response = _Response()._init_childptr(new Response(moveResponse(r)), None) - - func_obj = func - ret = func_obj(response) - Py_INCREF(ret) - return ret - + return func(response) cdef _find_field_order(struct_node): return [f.name for f in sorted(struct_node.fields, key=_attrgetter('codeOrder'))] @@ -84,7 +82,7 @@ cdef api VoidPromise * call_server_method(PyObject * _server, if type(ret) is _VoidPromise: return new VoidPromise(moveVoidPromise(deref((<_VoidPromise>ret).thisptr))) elif type(ret) is _Promise: - return new VoidPromise(helpers.convert_to_voidpromise(deref((<_Promise>ret).thisptr))) + return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) else: try: warning_msg = ( @@ -97,9 +95,9 @@ cdef api VoidPromise * call_server_method(PyObject * _server, if ret is not None: if type(ret) is _Promise: - return new VoidPromise(helpers.convert_to_voidpromise(deref((<_Promise>ret).thisptr))) + return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) elif type(ret) is _Promise: - return new VoidPromise(helpers.convert_to_voidpromise(deref((<_Promise>ret).thisptr))) + return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) else: try: warning_msg = ( @@ -120,7 +118,7 @@ cdef api VoidPromise * call_server_method(PyObject * _server, if type(ret) is _VoidPromise: return new VoidPromise(moveVoidPromise(deref((<_VoidPromise>ret).thisptr))) elif type(ret) is _Promise: - return new VoidPromise(helpers.convert_to_voidpromise(deref((<_Promise>ret).thisptr))) + return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) if not isinstance(ret, tuple): ret = (ret,) names = _find_field_order(context.results.schema.node.struct) @@ -136,30 +134,21 @@ cdef api VoidPromise * call_server_method(PyObject * _server, return NULL -cdef api convert_array_pyobject(PyArray & arr) with gil: - return [arr[i] for i in range(arr.size())] +cdef api object convert_array_pyobject(PyArray & arr) with gil: + return [arr[i].get().obj for i in range(arr.size())] -cdef api PyPromise * extract_promise(object obj) with gil: +cdef api Own[Promise[Own[PyRefCounter]]] extract_promise(object obj): if type(obj) is _Promise: - promise = <_Promise>obj - - ret = new PyPromise(promise.thisptr.attach(capnp.makePyRefCounter(promise))) - Py_DECREF(obj) - - return ret - - return NULL - - -cdef api RemotePromise * extract_remote_promise(object obj) with gil: - if type(obj) is _RemotePromise: - promise = <_RemotePromise>obj - promise.is_consumed = True - - return promise.thisptr # TODO:MEMORY: fix this leak - - return NULL + return move((<_Promise>obj).thisptr) + elif type(obj) is _RemotePromise: + parent = (<_RemotePromise>obj)._parent + # We don't need parent anymore. Setting to none allows quicker garbage collection + (<_RemotePromise>obj)._parent = None + return capnp.heap[PyPromise](helpers.convert_to_pypromise(move((<_RemotePromise>obj).thisptr)) + .attach(capnp.heap[PyRefCounter](parent))) + else: + return capnp.heap[PyPromise](capnp.heap[PyRefCounter](obj)) cdef extern from "" namespace " ::kj": @@ -317,7 +306,7 @@ ctypedef fused PromiseTypes: _Promise _RemotePromise _VoidPromise - PromiseFulfillerPair + # PromiseFulfillerPair cdef extern from "Python.h": @@ -1771,51 +1760,150 @@ cdef class _DynamicObjectBuilder: cpdef as_reader(self): return _DynamicObjectReader()._init(self.thisptr.asReader(), self._parent) +cdef void kjloop_runnable_callback(void* data) with gil: + cdef AsyncIoEventPort *port = data + assert port.runHandle is not None + port.timerImpl.advanceTo(systemPreciseMonotonicClock().now()) + port.kjLoop.run() + +cdef void kjloop_advance_callback(void* data) with gil: + cdef AsyncIoEventPort *port = data + assert port.runHandle is not None + port.timerImpl.advanceTo(systemPreciseMonotonicClock().now()) + +cdef cppclass AsyncIoEventPort(EventPort): + EventLoop *kjLoop + TimerImpl *timerImpl; + object asyncioLoop; + object runHandle; + + __init__(object asyncioLoop): + this.kjLoop = new EventLoop(deref(this)) + this.timerImpl = new TimerImpl(systemPreciseMonotonicClock().now()) + this.runHandle = None + this.asyncioLoop = asyncioLoop + + __dealloc__(): + del this.timerImpl + del this.kjLoop + + cbool wait() except* with gil: + raise KjException("Currently you cannot wait for promises while pycapnp is running in asyncio mode. " + + "You should instead use 'await'. If you have a use-case to start the asyncio loop " + + "using wait(), please report") + + cbool poll() except* with gil: + raise KjException("Currently you cannot poll promises while pycapnp is running in asyncio mode. " + + "If you have a use-case to poll the asyncio loop using poll(), please report") + + void setRunnable(cbool runnable) except* with gil: + if runnable: + if this.runHandle is not None: + # If a timer was running, cancel it and schedule a run immediately + # The timer will be re-scheduled once the kj loop becomes un-runnable again. + this.runHandle.cancel() + us = this; + this.runHandle = this.asyncioLoop.call_soon(lambda: kjloop_runnable_callback(us)) + else: + assert this.runHandle is not None + this.runHandle.cancel() + this.scheduleAdvance() + + void scheduleAdvance() with gil: + cdef uint64_t nextEvent = this.timerImpl.timeoutToNextEvent( + systemPreciseMonotonicClock().now(), MILLISECONDS, -1).orDefault(-1) + if nextEvent == -1: + this.runHandle = None + else: + seconds = nextEvent / 1000 + us = this; + this.runHandle = this.asyncioLoop.call_later(seconds, lambda: kjloop_advance_callback(us)) + + EventLoop *getKjLoop(): + return this.kjLoop + + Timer *getTimer(): + return this.timerImpl; + +def _asyncio_close_patch(loop, oldclose, _EventLoop kjloop): + # The purpose of patching the asyncio close() function is to set up the kj-loop to be closed as well. + # We replace the event loop getter with a weakref, such that it can be destroyed when all other + # references to it are gone. Then, if a new asyncio loop ever gets started, a new kj-loop can also be + # started. + _C_DEFAULT_EVENT_LOOP_LOCAL.loop = _weakref.ref(kjloop) + loop.close = oldclose() + return oldclose() cdef class _EventLoop: - cdef capnp.AsyncIoContext * thisptr + cdef object __weakref__ # Needed to make this class weak-referenceable + cdef Own[LowLevelAsyncIoProvider] lowLevelProvider + cdef Own[AsyncIoProvider] provider + cdef WaitScope * waitScope + cdef Timer* timer + cdef readonly in_asyncio_mode + + cdef AsyncIoEventPort *customPort def __init__(self): self._init() cdef _init(self) except +reraise_kj_exception: - self.thisptr = new capnp.AsyncIoContext(capnp.setupAsyncIo()) + try: + loop = asyncio.get_running_loop() + self.customPort = new AsyncIoEventPort(loop) + kjLoop = self.customPort.getKjLoop() + self.waitScope = new WaitScope(deref(kjLoop)) + self.timer = self.customPort.getTimer() + loop.close = _partial(_asyncio_close_patch, loop, loop.close, self) + self.in_asyncio_mode = True + except RuntimeError: + ptr = new capnp.AsyncIoContext(capnp.setupAsyncIo()) + self.lowLevelProvider = move(ptr.lowLevelProvider) + self.provider = move(ptr.provider) + self.waitScope = &ptr.waitScope + self.timer = &self.lowLevelProvider.get().getTimer() + del ptr + self.in_asyncio_mode = False def __dealloc__(self): - del self.thisptr #TODO:MEMORY: fix problems with Promises still being around - - cpdef _remove(self) except +reraise_kj_exception: - del self.thisptr - self.thisptr = NULL + if not self.customPort == NULL: + # If we have a custom port, the waitscope is not owned by provider, we have to delete it manually + del self.waitScope + del self.customPort cdef TwoWayPipe makeTwoWayPipe(self): - return deref(deref(self.thisptr).provider).newTwoWayPipe() + if self.in_asyncio_mode: + raise RuntimeError("Cannot call makeTwoWayPipe in asyncio mode") + return deref(self.provider).newTwoWayPipe() cdef Own[AsyncIoStream] wrapSocketFd(self, int fd): - return deref(deref(self.thisptr).lowLevelProvider).wrapSocketFd(fd) + if self.in_asyncio_mode: + raise RuntimeError("Cannot call wrapSocketFd in asyncio mode") + return deref(self.lowLevelProvider).wrapSocketFd(fd) -cdef _EventLoop C_DEFAULT_EVENT_LOOP - -_C_DEFAULT_EVENT_LOOP_LOCAL = None -_THREAD_LOCAL_EVENT_LOOPS = [] +_C_DEFAULT_EVENT_LOOP_LOCAL = _threading.local() cdef _EventLoop C_DEFAULT_EVENT_LOOP_GETTER(): - 'Optimization for not having to deal with threadlocal event loops unless we need to' - global C_DEFAULT_EVENT_LOOP - if C_DEFAULT_EVENT_LOOP is not None: - return C_DEFAULT_EVENT_LOOP - elif _C_DEFAULT_EVENT_LOOP_LOCAL is not None: - loop = getattr(_C_DEFAULT_EVENT_LOOP_LOCAL, 'loop', None) + global C_DEFAULT_EVENT_LOOP_LOCAL + loop = getattr(_C_DEFAULT_EVENT_LOOP_LOCAL, 'loop', None) + if type(loop) is _EventLoop: + return loop + elif type(loop) is _weakref.ref: + loop = loop() if loop is not None: - return <_EventLoop>_C_DEFAULT_EVENT_LOOP_LOCAL.loop + raise RuntimeError( + "The capnproto event loop associated to an already closed Python asyncio event loop is " + + "still running, because not all I/O events associated to it have terminated. If you wish " + + " to start a new loop, make sure that all previous events are cleaned up.") else: _C_DEFAULT_EVENT_LOOP_LOCAL.loop = _EventLoop() return _C_DEFAULT_EVENT_LOOP_LOCAL.loop else: - C_DEFAULT_EVENT_LOOP = _EventLoop() - return C_DEFAULT_EVENT_LOOP + assert loop is None + _C_DEFAULT_EVENT_LOOP_LOCAL.loop = _EventLoop() + return _C_DEFAULT_EVENT_LOOP_LOCAL.loop cdef class _Timer: @@ -1833,62 +1921,17 @@ def getTimer(): """ Get libcapnp event loop timer """ - return _Timer()._init(helpers.getTimer(C_DEFAULT_EVENT_LOOP_GETTER().thisptr)) + return _Timer()._init(C_DEFAULT_EVENT_LOOP_GETTER().timer) -cpdef remove_event_loop(ignore_errors=False): - '''Remove the global event loop''' - global C_DEFAULT_EVENT_LOOP - global _THREAD_LOCAL_EVENT_LOOPS +cpdef remove_event_loop(): + '''Remove the event loop''' global _C_DEFAULT_EVENT_LOOP_LOCAL - if C_DEFAULT_EVENT_LOOP: - try: - C_DEFAULT_EVENT_LOOP._remove() - except Exception as e: - if isinstance(ignore_errors, Exception): - if isinstance(e, ignore_errors): - ignore_errors = True - if ignore_errors is True: - pass - else: - raise - C_DEFAULT_EVENT_LOOP = None - if len(_THREAD_LOCAL_EVENT_LOOPS) > 0: - for loop in _THREAD_LOCAL_EVENT_LOOPS: - try: - loop._remove() - except Exception as e: - if isinstance(ignore_errors, Exception): - if isinstance(e, ignore_errors): - ignore_errors = True - if ignore_errors is True: - pass - else: - raise - _THREAD_LOCAL_EVENT_LOOPS = [] - _C_DEFAULT_EVENT_LOOP_LOCAL = None - - -cpdef create_event_loop(threaded=True): - '''Create a new global event loop. This will not remove the previous - EventLoop for you, so make sure to do that first''' - global C_DEFAULT_EVENT_LOOP - global _C_DEFAULT_EVENT_LOOP_LOCAL - if threaded: - if _C_DEFAULT_EVENT_LOOP_LOCAL is None: - _C_DEFAULT_EVENT_LOOP_LOCAL = _threading.local() - loop = _EventLoop() - _C_DEFAULT_EVENT_LOOP_LOCAL.loop = loop - _THREAD_LOCAL_EVENT_LOOPS.append(loop) - else: - C_DEFAULT_EVENT_LOOP = _EventLoop() - - -cpdef reset_event_loop(ignore_errors=False, threaded=True): - '''Removes event loop, then creates a new one. See remove_event_loop and create_event_loop for more details.''' - remove_event_loop(ignore_errors) - create_event_loop(threaded) + loop = getattr(_C_DEFAULT_EVENT_LOOP_LOCAL, 'loop', None) + if loop is not None: + loop._remove() + del _C_DEFAULT_EVENT_LOOP_LOCAL.loop def wait_forever(): @@ -1896,7 +1939,8 @@ def wait_forever(): Use libcapnp event loop to poll/wait forever """ cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() - helpers.waitNeverDone(deref(loop.thisptr).waitScope) + with nogil: + helpers.waitNeverDone(deref(loop.waitScope)) def poll_once(): @@ -1904,7 +1948,8 @@ def poll_once(): Poll libcapnp event loop once """ cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() - helpers.pollWaitScope(deref(loop.thisptr).waitScope) + with nogil: + loop.waitScope.poll() cdef class _CallContext: @@ -1936,194 +1981,64 @@ cdef class _CallContext: cpdef tail_call(self, _Request tailRequest): promise = _VoidPromise()._init(self.thisptr.tailCall(moveRequest(deref(tailRequest.thisptr_child)))) - promise.is_consumed = True return promise +cdef void _promise_check_consumed(PromiseTypes promise) except*: + if promise.thisptr.get() == NULL: + raise KjException( + "Promise was already used in a consuming operation. You can no longer use this Promise object") + +cdef _promise_then(PromiseTypes self, func, error_func, num_args, attach=None) except +reraise_kj_exception: + _promise_check_consumed(self) + + argspec = None + try: + argspec = _inspect.getfullargspec(func) + except (TypeError, ValueError): + pass + if argspec: + args_length = len(argspec.args) if argspec.args else 0 + defaults_length = len(argspec.defaults) if argspec.defaults else 0 + if args_length - defaults_length != num_args: + raise KjException(f'Function passed to `then` call must take exactly {num_args} arguments') + + return _Promise()._init( + helpers.then(move(self.thisptr), capnp.heap[PyRefCounter](func), + capnp.heap[PyRefCounter](error_func)) + .attach(capnp.heap[PyRefCounter]( attach))) + +cdef _promise_to_asyncio(PromiseTypes promise): + _promise_check_consumed(promise) + + fut = asyncio.get_running_loop().create_future() + # Attach the promise to the future, so that it doesn't get destroyed + fut.kjpromise = promise.then( + lambda res: fut.set_result(res) if not fut.cancelled() else None, + lambda err: fut.set_exception(err) if not fut.cancelled() else None) + del promise + fut.add_done_callback( + lambda fut: fut.kjpromise.cancel() if fut.cancelled() else None) + return fut cdef class _Promise: - cdef PyPromise * thisptr - cdef public bint is_consumed - cdef public object _parent, _obj - cdef _EventLoop _event_loop + cdef Own[PyPromise] thisptr def __init__(self, obj=None): - if obj is None: - self.is_consumed = True - else: - self.is_consumed = False - self._obj = obj - Py_INCREF(obj) # TODO: MEM: fix leak - self.thisptr = new PyPromise(obj) + if obj is not None: + self.thisptr = capnp.heap[PyPromise](capnp.heap[PyRefCounter](obj)) - self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() - - cdef _init(self, PyPromise other, parent=None): - self.is_consumed = False - self.thisptr = new PyPromise(movePromise(other)) - self._parent = parent + cdef _init(self, PyPromise other): + self.thisptr = capnp.heap[PyPromise](movePromise(other)) return self - def __dealloc__(self): - del self.thisptr - cpdef wait(self) except +reraise_kj_exception: - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - ret = helpers.waitPyPromise(self.thisptr, deref(self._event_loop.thisptr).waitScope) - Py_DECREF(ret) - - self.is_consumed = True - - return ret - - cpdef then(self, func, error_func=None) except +reraise_kj_exception: - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - argspec = None - try: - argspec = _inspect.getfullargspec(func) - except (TypeError, ValueError): - pass - if argspec: - args_length = len(argspec.args) if argspec.args else 0 - defaults_length = len(argspec.defaults) if argspec.defaults else 0 - if args_length - defaults_length != 1: - raise KjException('Function passed to `then` call must take exactly one argument') - - cdef _Promise new_promise = _Promise()._init( - helpers.then(deref(self.thisptr), func, error_func), self) - return _Promise()._init(new_promise.thisptr.attach( - capnp.makePyRefCounter(func), capnp.makePyRefCounter(error_func)), new_promise) - - def attach(self, *args): - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - ret = _Promise()._init(self.thisptr.attach(capnp.makePyRefCounter(args)), self) - self.is_consumed = True - - return ret - - cpdef cancel(self, numParents=1) except +reraise_kj_exception: - if numParents > 0 and hasattr(self._parent, 'cancel'): - self._parent.cancel(numParents - 1) - - self.is_consumed = True - del self.thisptr - self.thisptr = NULL - - -cdef class _VoidPromise: - cdef VoidPromise * thisptr - cdef public bint is_consumed - cdef public object _parent - cdef _EventLoop _event_loop - - def __init__(self): - self.is_consumed = True - self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() - - cdef _init(self, VoidPromise other, parent=None): - self.is_consumed = False - self.thisptr = new VoidPromise(moveVoidPromise(other)) - self._parent = parent - return self - - def __dealloc__(self): - del self.thisptr - - cpdef wait(self) except +reraise_kj_exception: - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - helpers.waitVoidPromise(self.thisptr, deref(self._event_loop.thisptr).waitScope) - - self.is_consumed = True - - cpdef then(self, func, error_func=None) except +reraise_kj_exception: - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - argspec = None - try: - argspec = _inspect.getfullargspec(func) - except (TypeError, ValueError): - pass - if argspec: - args_length = len(argspec.args) if argspec.args else 0 - defaults_length = len(argspec.defaults) if argspec.defaults else 0 - if args_length - defaults_length != 0: - raise KjException('Function passed to `then` call must take no arguments') - - cdef _Promise new_promise = _Promise()._init( - helpers.then(deref(self.thisptr), func, error_func), self) - return _Promise()._init(new_promise.thisptr.attach( - capnp.makePyRefCounter(func), capnp.makePyRefCounter(error_func)), new_promise) - - cpdef as_pypromise(self) except +reraise_kj_exception: - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - return _Promise()._init(helpers.convert_to_pypromise(deref(self.thisptr)), self) - - def attach(self, *args): - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - ret = _VoidPromise()._init(self.thisptr.attach(capnp.makePyRefCounter(args)), self) - self.is_consumed = True - - return ret - - cpdef cancel(self, numParents=1) except +reraise_kj_exception: - if numParents > 0 and hasattr(self._parent, 'cancel'): - self._parent.cancel(numParents - 1) - - self.is_consumed = True - del self.thisptr - self.thisptr = NULL - - -cdef class _RemotePromise: - cdef RemotePromise * thisptr - cdef public bint is_consumed - cdef public object _parent - cdef _EventLoop _event_loop - - def __init__(self): - self.is_consumed = True - self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() - - cdef _init(self, RemotePromise other, parent): - self.is_consumed = False - self.thisptr = new RemotePromise(moveRemotePromise(other)) - self._parent = parent - return self - - def __dealloc__(self): - del self.thisptr - - cpdef _wait(self) except +reraise_kj_exception: - return _Response()._init_childptr( - helpers.waitRemote(self.thisptr, deref(self._event_loop.thisptr).waitScope), self._parent) - - def wait(self): - """Wait on the promise. This will block until the promise has completed.""" - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - - ret = self._wait() - self.is_consumed = True - return ret + _promise_check_consumed(self) + cdef Own[PyPromise] prom = move(self.thisptr) # Explicit move to not leave thisptr dangling + cdef Own[PyRefCounter] ret + cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() + with nogil: + ret = move(prom.get().wait(deref(loop.waitScope))) + return ret.get().obj async def a_wait(self): """ @@ -2132,54 +2047,109 @@ cdef class _RemotePromise: Will still work with non-asyncio socket communication, but requires async handling of the function call. """ - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") + return await _promise_to_asyncio(self) - while not helpers.pollRemote(self.thisptr, deref(self._event_loop.thisptr).waitScope): - await asyncio.sleep(0.01) - ret = self._wait() - self.is_consumed = True - - return ret - - cpdef as_pypromise(self) except +reraise_kj_exception: - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") - return _Promise()._init(helpers.convert_to_pypromise(deref(self.thisptr)), self) + def __await__(self): + return _promise_to_asyncio(self).__await__() cpdef then(self, func, error_func=None) except +reraise_kj_exception: - """Promise pipelining, use to queue up operations on a promise before executing the promise.""" - if self.is_consumed: - raise KjException( - "Promise was already used in a consuming operation. You can no longer use this Promise object") + return _promise_then(self, func, error_func, 1) - argspec = None - try: - argspec = _inspect.getfullargspec(func) - except (TypeError, ValueError): - pass - if argspec: - args_length = len(argspec.args) if argspec.args else 0 - defaults_length = len(argspec.defaults) if argspec.defaults else 0 - if args_length - defaults_length != 1: - raise KjException('Function passed to `then` call must take exactly one argument') + cpdef cancel(self) except +reraise_kj_exception: + self.thisptr = Own[PyPromise]() - cdef _Promise new_promise = _Promise()._init( - helpers.then(deref(self.thisptr), func, error_func), self) - return _Promise()._init(new_promise.thisptr.attach( - capnp.makePyRefCounter(func), capnp.makePyRefCounter(error_func)), new_promise) + +cdef class _VoidPromise: + cdef Own[VoidPromise] thisptr + + + cdef _init(self, VoidPromise other): + self.thisptr = capnp.heap[VoidPromise](moveVoidPromise(other)) + return self + + cpdef wait(self) except +reraise_kj_exception: + _promise_check_consumed(self) + cdef Own[VoidPromise] prom = move(self.thisptr) # Explicit move to not leave thisptr dangling + cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() + with nogil: + prom.get().wait(deref(loop.waitScope)) + + async def a_wait(self): + """ + Asyncio version of wait(). + Required when using asyncio for socket communication. + + Will still work with non-asyncio socket communication, but requires async handling of the function call. + """ + # TODO: Is keeping a separate _VoidPromise class really worth it? Does it make things faster? + return await _promise_to_asyncio[_Promise](self.as_pypromise()) + + def __await__(self): + return _promise_to_asyncio[_Promise](self.as_pypromise()).__await__() + + cpdef as_pypromise(self) except +reraise_kj_exception: + _promise_check_consumed(self) + return _Promise()._init(helpers.convert_to_pypromise(move(self.thisptr))) + + cpdef then(self, func, error_func=None) except +reraise_kj_exception: + return _promise_then(self, func, error_func, 0) + + cpdef cancel(self) except +reraise_kj_exception: + self.thisptr = Own[VoidPromise]() + + + +cdef class _RemotePromise: + cdef object _parent + """A pointer to a parent object that needs to be kept alive for this promise to function. + Note that _Promise and _VoidPromise don't have such pointer. The reason is that in _RemotePromise + the parent pointer needs to be passed around through _RemotePromise._get. If an object needs to + be kept alive in _Promise or _VoidPromise, it can be attached to the underlying C++ promise.""" + + cdef Own[RemotePromise] thisptr + + cdef _init(self, RemotePromise other, object parent=None): + self.thisptr = capnp.heap[RemotePromise](moveRemotePromise(other)) + self._parent = parent + return self + + cpdef wait(self) except +reraise_kj_exception: + """Wait on the promise. This will block until the promise has completed.""" + _promise_check_consumed(self) + cdef _EventLoop loop = C_DEFAULT_EVENT_LOOP_GETTER() + with nogil: + response = helpers.waitRemote(move(self.thisptr), deref(loop.waitScope)) + return _Response()._init_childptr(response, None) + + async def a_wait(self): + """ + Asyncio version of wait(). + Required when using asyncio for socket communication. + + Will still work with non-asyncio socket communication, but requires async handling of the function call. + """ + return await _promise_to_asyncio(self) + + def __await__(self): + return _promise_to_asyncio(self).__await__() + + cpdef as_pypromise(self) except +reraise_kj_exception: + _promise_check_consumed(self) + parent = self._parent + self._parent = None # We don't need parent anymore. Setting to none allows quicker garbage collection + return _Promise()._init(helpers.convert_to_pypromise(move(self.thisptr)) + .attach(capnp.heap[PyRefCounter](parent))) cpdef _get(self, field) except +reraise_kj_exception: - cdef int type = (self.thisptr.get(field)).getType() + _promise_check_consumed(self) + cdef int type = (self.thisptr.get().get(field)).getType() if type == capnp.TYPE_CAPABILITY: return _DynamicCapabilityClient()._init( - (self.thisptr.get(field)).asCapability(), self._parent) + (self.thisptr.get().get(field)).asCapability(), self._parent) elif type == capnp.TYPE_STRUCT: return _DynamicStructPipeline()._init( new C_DynamicStruct.Pipeline( - (self.thisptr.get(field)).asStruct()), self._parent) + (self.thisptr.get().get(field)).asStruct()), self._parent) elif type == capnp.TYPE_UNKNOWN: raise KjException("Cannot convert type to Python. Type is unknown by capnproto library") else: @@ -2194,7 +2164,8 @@ cdef class _RemotePromise: property schema: """A property that returns the _StructSchema object matching this reader""" def __get__(self): - return _StructSchema()._init_child(self.thisptr.getSchema()) + _promise_check_consumed(self) + return _StructSchema()._init_child(self.thisptr.get().getSchema()) def __dir__(self): return list(set(self.schema.fieldnames + tuple(dir(self.__class__)))) @@ -2202,23 +2173,14 @@ cdef class _RemotePromise: def to_dict(self, verbose=False, ordered=False): return _to_dict(self, verbose, ordered) - cpdef cancel(self, numParents=1) except +reraise_kj_exception: - if numParents > 0 and hasattr(self._parent, 'cancel'): - self._parent.cancel(numParents - 1) + cpdef then(self, func, error_func=None) except +reraise_kj_exception: + parent = self._parent + self._parent = None # We don't need parent anymore. Setting to none allows quicker garbage collection + return _promise_then(self, func, error_func, 1, attach=parent) - self.is_consumed = True - del self.thisptr - self.thisptr = NULL - - # def attach(self, *args): - # if self.is_consumed: - # raise KjException( - # "Promise was already used in a consuming operation. You can no longer use this Promise object") - - # ret = _RemotePromise()._init(self.thisptr.attach(capnp.makePyRefCounter(args)), self) - # self.is_consumed = True - - # return ret + cpdef cancel(self) except +reraise_kj_exception: + self.thisptr = Own[RemotePromise]() + self._parent = None # We don't need parent anymore. Setting to none allows quicker garbage collection cpdef join_promises(promises) except +reraise_kj_exception: @@ -2238,7 +2200,6 @@ cpdef join_promises(promises) except +reraise_kj_exception: raise KjException( "One of the promises passed to `join_promises` had a non promise value of: {}".format(promise)) heap.add(movePromise(deref(pyPromise.thisptr))) - pyPromise.is_consumed = True return _Promise()._init(helpers.then(capnp.joinPromises(heap.finish()))) @@ -2302,7 +2263,7 @@ cdef class _DynamicCapabilityServer: cdef class _DynamicCapabilityClient: cdef C_DynamicCapability.Client thisptr - cdef public object _server, _parent, _cached_schema + cdef public object _parent, _cached_schema cdef _init(self, C_DynamicCapability.Client other, object parent): self.thisptr = other @@ -2317,7 +2278,7 @@ cdef class _DynamicCapabilityClient: s = schema self.thisptr = helpers.new_client(s.thisptr, server) - self._server = server + self._parent = server return self cpdef _find_method_args(self, method_name): @@ -2449,7 +2410,7 @@ cdef class _TwoPartyVatNetwork: return self cpdef on_disconnect(self) except +reraise_kj_exception: - return _VoidPromise()._init(deref(self.thisptr).onDisconnect(), self) + return _VoidPromise()._init(deref(self.thisptr).onDisconnect()) cdef class TwoPartyClient: @@ -2465,32 +2426,29 @@ cdef class TwoPartyClient: """ cdef RpcSystem * thisptr cdef public _TwoPartyVatNetwork _network - cdef public object _orig_stream - cdef public _AsyncIoStream _stream cdef public _TwoWayPipe _pipe def __init__(self, socket=None, traversal_limit_in_words=None, nesting_limit=None): if isinstance(socket, basestring): + if C_DEFAULT_EVENT_LOOP_GETTER().in_asyncio_mode: + raise RuntimeError("Pycapnp is in asyncio mode. Pass a AsyncIoStream") socket = self._connect(socket) cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit) - self._orig_stream = socket - if self._orig_stream: - self._stream = _FdAsyncIoStream(socket.fileno()) - self._network = _TwoPartyVatNetwork()._init(self._stream, capnp.CLIENT, opts) - else: + if socket is None: # Initialize TwoWayPipe, to use pipe() acquire other end of the pipe using read() and write() methods self._pipe = _TwoWayPipe() self._network = _TwoPartyVatNetwork()._init_pipe(self._pipe, capnp.CLIENT, opts) + elif isinstance(socket, _AsyncIoStream): + self._network = _TwoPartyVatNetwork()._init(socket, capnp.CLIENT, opts) + elif isinstance(socket, _socket.socket): + stream = _FdAsyncIoStream(socket) + self._network = _TwoPartyVatNetwork()._init(stream, capnp.CLIENT, opts) + else: + raise ValueError(f"Argument socket should be a string, socket, AsyncIoStream or None, was {type(socket)}") self.thisptr = new RpcSystem(makeRpcClient(deref(self._network.thisptr))) - if self._orig_stream: - Py_INCREF(self._orig_stream) - Py_INCREF(self._stream) - else: - Py_INCREF(self._pipe) - Py_INCREF(self._network) # TODO:MEMORY: attach this to onDrained, also figure out what's leaking async def read(self, bufsize): """ @@ -2498,45 +2456,40 @@ cdef class TwoPartyClient: :param bufsize: Buffer size to read from the libcapnp library """ - cdef AsyncIoStreamReadHelper *reader = new AsyncIoStreamReadHelper( - self._pipe._pipe.ends[1].get(), - &self._pipe._event_loop.thisptr.waitScope, - bufsize - ) - while not reader.poll(): - await asyncio.sleep(0.01) cdef array.array read_buffer = array.array('b', []) - array.resize(read_buffer, reader.read_size()) - memcpy(read_buffer.data.as_voidptr, reader.read_buffer(), reader.read_size()) - del reader + array.resize(read_buffer, bufsize) + read_size_actual = await _Promise()._init( + helpers.wrapSizePromise( + self._pipe._pipe.ends[1].get().read(read_buffer.data.as_voidptr, 1, bufsize))) + array.resize(read_buffer, read_size_actual) return read_buffer - def write(self, data): + async def write(self, data): """ libcapnp writer (asyncio sockets only) :param data: Buffer to write to the libcapnp library """ cdef array.array write_buffer = array.array('b', data) - deref(self._pipe._pipe.ends[1]).write( - write_buffer.data.as_voidptr, - len(data) - ).wait(self._pipe._event_loop.thisptr.waitScope) + await _VoidPromise()._init( + deref(self._pipe._pipe.ends[1]).write( + write_buffer.data.as_voidptr, + len(data) + )) def __dealloc__(self): - del self.thisptr + if not self.thisptr == NULL: + del self.thisptr - cpdef _connect(self, host_string): + def _connect(self, host_string): if host_string.startswith('unix:'): path = host_string[5:] sock = _socket.socket(_socket.AF_UNIX, _socket.SOCK_STREAM) sock.connect(path) else: host, port = host_string.split(':') - sock = _socket.create_connection((host, port)) - # Set TCP_NODELAY on socket to disable Nagle's algorithm. This is not # neccessary, but it speeds things up. sock.setsockopt(_socket.IPPROTO_TCP, _socket.TCP_NODELAY, 1) @@ -2564,8 +2517,6 @@ cdef class TwoPartyServer: """ cdef RpcSystem * thisptr cdef public _TwoPartyVatNetwork _network - cdef public object _orig_stream, _disconnect_promise - cdef public _AsyncIoStream _stream cdef public _TwoWayPipe _pipe cdef object _port cdef public object port_promise, _bootstrap @@ -2580,19 +2531,27 @@ cdef class TwoPartyServer: self._bootstrap = None if isinstance(socket, basestring): + if C_DEFAULT_EVENT_LOOP_GETTER().in_asyncio_mode: + raise RuntimeError("Pycapnp is in asyncio mode. Please start an asyncio server using" + "TwoPartyClient.create_server and pass any resulting connection to this class.") self._connect(socket, bootstrap, traversal_limit_in_words, nesting_limit) return - self._orig_stream = socket - if self._orig_stream: - self._stream = _FdAsyncIoStream(socket.fileno()) - self._network = _TwoPartyVatNetwork()._init( - self._stream, capnp.SERVER, make_reader_opts(traversal_limit_in_words, nesting_limit)) - else: + opts = make_reader_opts(traversal_limit_in_words, nesting_limit) + if isinstance(socket, _AsyncIoStream): + self._network = _TwoPartyVatNetwork()._init(socket, capnp.SERVER, opts) + elif isinstance(socket, _socket.socket): + if C_DEFAULT_EVENT_LOOP_GETTER().in_asyncio_mode: + raise RuntimeError("Pycapnp is in asyncio mode. Please pass an AsyncIoStream instance.") + stream = _FdAsyncIoStream(socket) + self._network = _TwoPartyVatNetwork()._init(stream, capnp.SERVER, opts) + elif socket is None: # Initialize TwoWayPipe, to use pipe() acquire other end of the pipe using read() and write() methods self._pipe = _TwoWayPipe() self._network = _TwoPartyVatNetwork()._init_pipe( - self._pipe, capnp.SERVER, make_reader_opts(traversal_limit_in_words, nesting_limit)) + self._pipe, capnp.SERVER, opts) + else: + raise KjException("Unexpected typ for socket in TwoPartyServer") self._port = 0 @@ -2602,31 +2561,18 @@ cdef class TwoPartyServer: self.thisptr = new RpcSystem(makeRpcServerBootstrap( deref(self._network.thisptr), helpers.server_to_client(schema.thisptr, bootstrap))) - Py_INCREF(self._orig_stream) - Py_INCREF(self._stream) - Py_INCREF(self._pipe) - Py_INCREF(self._bootstrap) - Py_INCREF(self._network) - self._disconnect_promise = self.on_disconnect().then(self._decref) - async def read(self, bufsize): """ libcapnp reader (asyncio sockets only) :param bufsize: Buffer size to read from the libcapnp library """ - cdef AsyncIoStreamReadHelper *reader = new AsyncIoStreamReadHelper( - self._pipe._pipe.ends[1].get(), - &self._pipe._event_loop.thisptr.waitScope, - bufsize - ) - while not reader.poll(): - await asyncio.sleep(0.01) - cdef array.array read_buffer = array.array('b', []) - array.resize(read_buffer, reader.read_size()) - memcpy(read_buffer.data.as_voidptr, reader.read_buffer(), reader.read_size()) - del reader + array.resize(read_buffer, bufsize) + read_size_actual = await _Promise()._init( + helpers.wrapSizePromise( + self._pipe._pipe.ends[1].get().read(read_buffer.data.as_voidptr, 1, bufsize))) + array.resize(read_buffer, read_size_actual) return read_buffer async def write(self, data): @@ -2636,10 +2582,11 @@ cdef class TwoPartyServer: :param data: Buffer to write to the libcapnp library """ cdef array.array write_buffer = array.array('b', data) - deref(self._pipe._pipe.ends[1]).write( - write_buffer.data.as_voidptr, - len(data) - ).wait(self._pipe._event_loop.thisptr.waitScope) + await _VoidPromise()._init( + deref(self._pipe._pipe.ends[1]).write( + write_buffer.data.as_voidptr, + len(data) + )) cpdef _connect(self, host_string, bootstrap, traversal_limit_in_words, nesting_limit): cdef schema_cpp.ReaderOptions opts = make_reader_opts(traversal_limit_in_words, nesting_limit) @@ -2649,26 +2596,20 @@ cdef class TwoPartyServer: self._task_set = new capnp.TaskSet(self._error_handler) 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, opts)) - - def _decref(self): - Py_DECREF(self._bootstrap) - Py_INCREF(self._pipe) - Py_DECREF(self._orig_stream) - Py_DECREF(self._stream) - Py_DECREF(self._network) + loop.provider.get(), temp_string, opts)) def __dealloc__(self): del self.thisptr del self._task_set cpdef on_disconnect(self) except +reraise_kj_exception: + if self._task_set != NULL: + raise KjException("Currently, you can only call on_disconnect on a server without an internal socket") return _VoidPromise()._init(deref(self._network.thisptr).onDisconnect()) def poll_once(self): @@ -2678,12 +2619,11 @@ cdef class TwoPartyServer: return poll_once() async def poll_forever(self): - """ + """Deprecated. Do not use. Poll libcapnp library forever (asyncio) """ - while True: - poll_once() - await asyncio.sleep(0.01) + raise KjException("This functionality has been removed. If you wish to wait forever, use \n" + + "'await asyncio._get_running_loop().create_future()'") cpdef run_forever(self): if self.port_promise is None: @@ -2705,7 +2645,275 @@ cdef class TwoPartyServer: cdef class _AsyncIoStream: cdef Own[AsyncIoStream] thisptr + cdef _EventLoop _event_loop # We hold a pointer to the event loop here, to ensure it remains alive + @staticmethod + async def create_connection(host = None, port = None, **kwargs): + """Create a TCP connection. + + All parameters given to this function are passed to `asyncio.get_running_loop().create_connection()`. + See that function for documentation on the possible arguments. + """ + cdef _AsyncIoStream self = _AsyncIoStream() + self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() + loop = asyncio.get_running_loop() + transport, protocol = await loop.create_connection( + lambda: _PyAsyncIoStreamProtocol(), host, port, **kwargs) + self.thisptr = capnp.heap[PyAsyncIoStream](capnp.heap[PyRefCounter](protocol)) + return self + + @staticmethod + async def create_unix_connection(path = None, **kwargs): + """Create a Unix socket connection. + + All parameters given to this function are passed to `asyncio.get_running_loop().create_unix_connection()`. + See that function for documentation on the possible arguments. + """ + cdef _AsyncIoStream self = _AsyncIoStream() + self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() + loop = asyncio.get_running_loop() + transport, protocol = await loop.create_unix_connection( + lambda: _PyAsyncIoStreamProtocol(), path, **kwargs) + self.thisptr = capnp.heap[PyAsyncIoStream](capnp.heap[PyRefCounter](protocol)) + return self + + @staticmethod + def _connect(callback): + cdef _AsyncIoStream self = _AsyncIoStream() + self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() + loop = asyncio.get_running_loop() + protocol = _PyAsyncIoStreamProtocol(callback, self) + self.thisptr = capnp.heap[PyAsyncIoStream](capnp.heap[PyRefCounter](protocol)) + return protocol + + @staticmethod + async def create_server(callback, host = None, port = None, **kwargs): + """Create a TCP connection server. + + The `callback` parameter will be called whenever a new connection is made. It receives a `AsyncIoStream` + instance as its only argument. If the result of `callback` is a coroutine, it will be scheduled as a task. + + This function behaves similarly to `asyncio.get_running_loop().create_server()`. All arguments except + for `callback` will be passed directly to that function, and the server returned is similar as well. + See that function for documentation on the possible arguments. + """ + loop = asyncio.get_running_loop() + return await loop.create_server(lambda: _AsyncIoStream._connect(callback), host, port, **kwargs) + + @staticmethod + async def create_unix_server(callback, path = None, **kwargs): + """Create a unix connection server. + + The `callback` parameter will be called whenever a new connection is made. It receives a `AsyncIoStream` + instance as its only argument. If the result of `callback` is a coroutine, it will be scheduled as a task. + + This function behaves similarly to `asyncio.get_running_loop().create_server()`. All arguments except + for `callback` will be passed directly to that function, and the server returned is similar as well. + See that function for documentation on the possible arguments. + """ + loop = asyncio.get_running_loop() + return await loop.create_unix_server(lambda: _AsyncIoStream._connect(callback), path, **kwargs) + +cdef class DummyBaseClass: + pass + +cdef class _PyAsyncIoStreamProtocol(DummyBaseClass, asyncio.BufferedProtocol): + # TODO: Temporary. Needed due to a missing __slots__ definitions in BufferedProtocol on Python 3.7. + # See https://github.com/python/cpython/issues/79575. Can be removed once Python 3.7 is unsupported. + cdef dict __dict__ + + cdef object transport + cdef object connected_callback + cdef object callback_arg + + # State for reading data from the transport + cdef char* read_buffer + cdef size_t read_min_bytes + cdef size_t read_max_bytes + cdef size_t read_already_read + cdef PromiseFulfiller[size_t]* read_fulfiller + cdef cbool read_eof + + # TODO: Temporary. This is an overflow buffer, which is needed for two blatant violations of the protocol. + # The first violation is int the SSL transport implementation. + # See https://github.com/python/cpython/issues/89322, fixed in Python 3.11. This bug causes the + # SSL transport to force data upon us even when we've asked it to pause sending us data. Therefore, + # we have to store the data in a overflow buffer. + # + # The second violation is that a transport cannot be paused immediately after it is connected. + # See https://github.com/python/cpython/issues/103607. This also causes the need to be prepared + # for unexpected data. + # + # This extra code can be removed once both bugs are fixed in all supported python versions. + cdef bytearray read_overflow_buffer + cdef bytearray read_overflow_buffer_current + + # State for writing data to the transport + cdef cbool write_paused + cdef cbool write_in_progress + cdef ArrayPtr[const ArrayPtr[const uint8_t]] write_pieces + cdef size_t write_index + cdef VoidPromiseFulfiller* write_fulfiller + + def __init__(self, connected_callback = None, callback_arg = None): + self.connected_callback = connected_callback + self.callback_arg = callback_arg + + def connection_made(self, transport): + self.transport = transport + + # TODO: BUG. We want to immediately pause reading, but Python's transport implementation does not + # allow this. See https://github.com/python/cpython/issues/103607. + # To work around this, we also insert pause_reading() in get_buffer() when appropriate. + transport.pause_reading() + + self.write_paused = False + self.write_in_progress = False + self.read_eof = False + self.read_overflow_buffer = bytearray() + if self.connected_callback is not None: + callback_res = self.connected_callback(self.callback_arg) + if asyncio.iscoroutine(callback_res): + asyncio.get_running_loop().create_task(callback_res) + self.connected_callback = None + self.callback_arg = None + + def connection_lost(self, exc): + if self.read_fulfiller != NULL: + capnp.rejectDisconnected[size_t](deref(self.read_fulfiller), StringPtr(str(exc))) + self.read_buffer = NULL + self.read_fulfiller = NULL + if self.write_fulfiller != NULL: + capnp.rejectVoidDisconnected(deref(self.write_fulfiller), StringPtr(str(exc))) + self.write_reset() + self.write_paused = True + self.transport = None + + def get_buffer(self, size_hint): + if self.read_buffer == NULL: # Should not happen, but for SSL it does, see comment above + + # TODO: Bug. Workaround for the transport ignoring pause_reading() in connection_made() + self.transport.pause_reading() + + size = size_hint if size_hint > 0 else 100 + self.read_overflow_buffer_current = bytearray(size) + return self.read_overflow_buffer_current + else: + return memoryview.PyMemoryView_FromMemory(self.read_buffer, self.read_max_bytes, buffer.PyBUF_WRITE) + + def buffer_updated(self, size): + if self.read_buffer == NULL: # Should not happen, but for SSL it does, see comment above + self.read_overflow_buffer.extend(self.read_overflow_buffer_current[0:size]) + else: + self.read_buffer += size + self.read_min_bytes -= size + self.read_max_bytes -= size + self.read_already_read += size + if self.read_min_bytes <= 0: + self.read_fulfiller.fulfill(move(self.read_already_read)) + self.read_reset() + + def pause_writing(self): + self.write_paused = True + + def resume_writing(self): + self.write_paused = False + self.write_loop() + + def eof_received(self): + self.read_eof = True + if self.read_buffer != NULL: + self.read_fulfiller.fulfill(move(self.read_already_read)) + self.read_reset() + + cdef write_loop(self): + if not self.write_in_progress: return + cdef const ArrayPtr[const uint8_t]* piece + for i in range(self.write_index, self.write_pieces.size()): + if self.write_paused: + self.write_index = i + break + piece = &self.write_pieces[i] + view = memoryview.PyMemoryView_FromMemory(piece.begin(), piece.size(), buffer.PyBUF_READ) + self.transport.write(view) + if not self.write_paused: + self.write_fulfiller.fulfill() + self.write_reset() + + cdef read_reset(self): + self.transport.pause_reading() + self.read_buffer = NULL + self.read_fulfiller = NULL + + cdef write_reset(self): + self.write_in_progress = False + self.write_fulfiller = NULL + + +cdef api void _asyncio_stream_write_start( + object thisptr, ArrayPtr[const ArrayPtr[const uint8_t]] pieces, + VoidPromiseFulfiller& fulfiller) except*: + cdef _PyAsyncIoStreamProtocol self = <_PyAsyncIoStreamProtocol>thisptr + if self.transport is None or self.transport.is_closing(): + capnp.rejectVoidDisconnected(fulfiller, StringPtr("Socket is closing.")) + return + self.write_pieces = pieces + self.write_index = 0 + self.write_fulfiller = &fulfiller + self.write_in_progress = True + self.write_loop() + +cdef api void _asyncio_stream_write_stop(object thisptr): + (<_PyAsyncIoStreamProtocol>thisptr).write_reset() + +cdef api void _asyncio_stream_read_start( + object thisptr, void* buffer, size_t min_bytes, size_t max_bytes, + PromiseFulfiller[size_t]& fulfiller) except*: + cdef _PyAsyncIoStreamProtocol self = <_PyAsyncIoStreamProtocol>thisptr + if self.transport is None or self.transport.is_closing(): + capnp.rejectDisconnected(fulfiller, StringPtr("Socket is closing")) + return + if self.read_eof: + self.read_fulfiller.fulfill(0) + return + self.read_buffer = buffer + self.read_min_bytes = min_bytes + self.read_max_bytes = max_bytes + self.read_already_read = 0 + self.read_fulfiller = &fulfiller + + # Begin of draining the overflow buffer, which is created because of a bug in SSL, see comment above. + # Can be removed once Python < 3.11 is not longer supported. + if self.read_overflow_buffer: + to_copy = min(len(self.read_overflow_buffer), max_bytes) + memcpy(buffer, self.read_overflow_buffer, to_copy) + del self.read_overflow_buffer[:to_copy] + self.read_buffer += to_copy + self.read_min_bytes -= to_copy + self.read_max_bytes -= to_copy + self.read_already_read += to_copy + if self.read_min_bytes <= 0: + self.read_fulfiller.fulfill(move(self.read_already_read)) + self.read_reset() + return # resume_reading no longer needed + # End of draining the overflow buffer. + + self.transport.resume_reading() + +cdef api void _asyncio_stream_read_stop(object thisptr): + cdef _PyAsyncIoStreamProtocol self = <_PyAsyncIoStreamProtocol>thisptr + if self.transport is not None: self.read_reset() + +cdef api void _asyncio_stream_shutdown_write(object thisptr) except*: + cdef _PyAsyncIoStreamProtocol self = <_PyAsyncIoStreamProtocol>thisptr + if self.transport is not None and self.transport.can_write_eof(): + self.transport.write_eof() + +cdef api void _asyncio_stream_close(object thisptr) except*: + cdef _PyAsyncIoStreamProtocol self = <_PyAsyncIoStreamProtocol>thisptr + # Careful, the transport object may have already been partially destroyed here. + if self.transport is not None and hasattr(self.transport, "close"): + self.transport.close() cdef class _TwoWayPipe: cdef _EventLoop _event_loop @@ -2721,19 +2929,24 @@ cdef class _TwoWayPipe: cdef class _FdAsyncIoStream(_AsyncIoStream): - cdef _EventLoop _event_loop + """Wraps a socket for usage with pycapnp. + Note that this class does not own the socket. Instead, it receives the fileno from a python object, + which will continue to own it. This object is kept alive as long as this class is alive. Ultimately, + the python object is responsible for closing.""" + cdef object _socket - def __init__(self, int fd): - self._init(fd) + def __init__(self, object socket): + self._socket = socket + self._init(socket.fileno()) cdef _init(self, int fd) except +reraise_kj_exception: self._event_loop = C_DEFAULT_EVENT_LOOP_GETTER() self.thisptr = self._event_loop.wrapSocketFd(fd) - -cdef class PyAsyncIoStream(_AsyncIoStream): - def __init__(self, int fd): - pass + def __dealloc__(self): + # The AsyncIoStream must be destroyed before self._socket is removed, to ensure the socket is still + # open when the destructor is called. Therefore, we do this manually to have control over the ordering + self.thisptr = Own[AsyncIoStream]() cdef class PromiseFulfillerPair: @@ -3413,9 +3626,11 @@ class _InterfaceModule(object): self.Server = type(name + '.Server', (_DynamicCapabilityServer,), {'__init__': server_init, 'schema':schema}) def _new_client(self, server): + C_DEFAULT_EVENT_LOOP_GETTER() # Make sure that the event loop has been initialized return _DynamicCapabilityClient()._init_vals(self.schema, server) def _new_server(self, server): + C_DEFAULT_EVENT_LOOP_GETTER() # Make sure that the event loop has been initialized return _DynamicCapabilityServer(self.schema, server) diff --git a/examples/async_calculator_client.py b/examples/async_calculator_client.py index 33940f7..41ba8f8 100755 --- a/examples/async_calculator_client.py +++ b/examples/async_calculator_client.py @@ -2,7 +2,6 @@ import argparse import asyncio -import socket import capnp import calculator_capnp @@ -24,19 +23,6 @@ class PowerFunction(calculator_capnp.Calculator.Function.Server): return pow(params[0], params[1]) -async def myreader(client, reader): - while True: - data = await reader.read(4096) - client.write(data) - - -async def mywriter(client, writer): - while True: - data = await client.read(4096) - writer.write(data.tobytes()) - await writer.drain() - - def parse_args(): parser = argparse.ArgumentParser( usage="Connects to the Calculator server \ @@ -48,27 +34,9 @@ at the given address and does some RPCs" async def main(host): - host = host.split(":") - addr = host[0] - port = host[1] - # Handle both IPv4 and IPv6 cases - try: - print("Try IPv4") - reader, writer = await asyncio.open_connection( - addr, port, family=socket.AF_INET - ) - except Exception: - print("Try IPv6") - reader, writer = await asyncio.open_connection( - addr, port, family=socket.AF_INET6 - ) - - # Start TwoPartyClient using TwoWayPipe (takes no arguments in this mode) - client = capnp.TwoPartyClient() - - # Assemble reader and writer tasks, run in the background - coroutines = [myreader(client, reader), mywriter(client, writer)] - asyncio.gather(*coroutines, return_exceptions=True) + host, port = parse_args().host.split(":") + connection = await capnp.AsyncIoStream.create_connection(host=host, port=port) + client = capnp.TwoPartyClient(connection) # Bootstrap the Calculator interface calculator = client.bootstrap().cast_as(calculator_capnp.Calculator) @@ -106,7 +74,7 @@ async def main(host): # Now that we've sent all the requests, wait for the response. Until this # point, we haven't waited at all! - response = await read_promise.a_wait() + response = await read_promise assert response.value == 123 print("PASS") @@ -144,7 +112,7 @@ async def main(host): eval_promise = request.send() read_promise = eval_promise.value.read() - response = await read_promise.a_wait() + response = await read_promise assert response.value == 101 print("PASS") @@ -208,8 +176,8 @@ async def main(host): add_5_promise = add_5_request.send().value.read() # Now wait for the results. - assert (await add_3_promise.a_wait()).value == 27 - assert (await add_5_promise.a_wait()).value == 29 + assert (await add_3_promise).value == 27 + assert (await add_5_promise).value == 29 print("PASS") @@ -290,8 +258,8 @@ async def main(host): g_eval_promise = g_eval_request.send().value.read() # Wait for the results. - assert (await f_eval_promise.a_wait()).value == 1234 - assert (await g_eval_promise.a_wait()).value == 4244 + assert (await f_eval_promise).value == 1234 + assert (await g_eval_promise).value == 4244 print("PASS") @@ -329,7 +297,7 @@ async def main(host): add_params[1].literal = 5 # Send the request and wait. - response = await request.send().value.read().a_wait() + response = await request.send().value.read() assert response.value == 512 print("PASS") diff --git a/examples/async_calculator_server.py b/examples/async_calculator_server.py index 339ff6f..f1c6dc4 100755 --- a/examples/async_calculator_server.py +++ b/examples/async_calculator_server.py @@ -3,7 +3,6 @@ import argparse import asyncio import logging -import socket import capnp import calculator_capnp @@ -13,60 +12,6 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.DEBUG) -class Server: - async def myreader(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.reader.read(4096), timeout=0.1) - except asyncio.TimeoutError: - logger.debug("myreader timeout.") - continue - except Exception as err: - logger.error("Unknown myreader err: %s", err) - return False - await self.server.write(data) - logger.debug("myreader done.") - return True - - async def mywriter(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.server.read(4096), timeout=0.1) - self.writer.write(data.tobytes()) - except asyncio.TimeoutError: - logger.debug("mywriter timeout.") - continue - except Exception as err: - logger.error("Unknown mywriter err: %s", err) - return False - logger.debug("mywriter done.") - return True - - async def myserver(self, reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - self.server = capnp.TwoPartyServer(bootstrap=CalculatorImpl()) - self.reader = reader - self.writer = writer - self.retry = True - - # Assemble reader and writer tasks, run in the background - coroutines = [self.myreader(), self.mywriter()] - tasks = asyncio.gather(*coroutines, return_exceptions=True) - - while True: - self.server.poll_once() - # Check to see if reader has been sent an eof (disconnect) - if self.reader.at_eof(): - self.retry = False - break - await asyncio.sleep(0.01) - - # Make wait for reader/writer to finish (prevent possible resource leaks) - await tasks - - def read_value(value): """Helper function to asynchronously call read() on a Calculator::Value and return a promise for the result. (In the future, the generated code might @@ -180,6 +125,11 @@ class CalculatorImpl(calculator_capnp.Calculator.Server): return OperatorImpl(op) +async def new_connection(stream): + server = capnp.TwoPartyServer(stream, bootstrap=CalculatorImpl()) + await server.on_disconnect() + + def parse_args(): parser = argparse.ArgumentParser( usage="""Runs the server bound to the\ @@ -191,29 +141,9 @@ given address/port ADDRESS. """ return parser.parse_args() -async def new_connection(reader, writer): - server = Server() - await server.myserver(reader, writer) - - async def main(): - address = parse_args().address - host = address.split(":") - addr = host[0] - port = host[1] - - # Handle both IPv4 and IPv6 cases - try: - print("Try IPv4") - server = await asyncio.start_server( - new_connection, addr, port, family=socket.AF_INET - ) - except Exception: - print("Try IPv6") - server = await asyncio.start_server( - new_connection, addr, port, family=socket.AF_INET6 - ) - + host, port = parse_args().address.split(":") + server = await capnp.AsyncIoStream.create_server(new_connection, host, port) async with server: await server.serve_forever() diff --git a/examples/async_client.py b/examples/async_client.py index 67a79b2..a97e212 100755 --- a/examples/async_client.py +++ b/examples/async_client.py @@ -4,13 +4,9 @@ import asyncio import argparse import time import capnp -import socket import thread_capnp -capnp.remove_event_loop() -capnp.create_event_loop(threaded=True) - def parse_args(): parser = argparse.ArgumentParser( @@ -29,61 +25,35 @@ class StatusSubscriber(thread_capnp.Example.StatusSubscriber.Server): print("status: {}".format(time.time())) -async def myreader(client, reader): - while True: - data = await reader.read(4096) - client.write(data) - - -async def mywriter(client, writer): - while True: - data = await client.read(4096) - writer.write(data.tobytes()) - - async def background(cap): subscriber = StatusSubscriber() - promise = cap.subscribeStatus(subscriber) - await promise.a_wait() + await cap.subscribeStatus(subscriber) async def main(host): - host = host.split(":") - addr = host[0] - port = host[1] - # Handle both IPv4 and IPv6 cases - try: - print("Try IPv4") - reader, writer = await asyncio.open_connection( - addr, port, family=socket.AF_INET - ) - except Exception: - print("Try IPv6") - reader, writer = await asyncio.open_connection( - addr, port, family=socket.AF_INET6 - ) - - # Start TwoPartyClient using TwoWayPipe (takes no arguments in this mode) - client = capnp.TwoPartyClient() + host, port = host.split(":") + connection = await capnp.AsyncIoStream.create_connection(host=host, port=port) + client = capnp.TwoPartyClient(connection) cap = client.bootstrap().cast_as(thread_capnp.Example) - # Assemble reader and writer tasks, run in the background - coroutines = [myreader(client, reader), mywriter(client, writer)] - asyncio.gather(*coroutines, return_exceptions=True) - # Start background task for subscriber - tasks = [background(cap)] - asyncio.gather(*tasks, return_exceptions=True) + asyncio.create_task(background(cap)) # Run blocking tasks print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) if __name__ == "__main__": - asyncio.run(main(parse_args().host)) + args = parse_args() + asyncio.run(main(args.host)) + + # Test that we can run multiple asyncio loops in sequence. This is particularly tricky, because + # main contains a background task that we never cancel. The entire loop gets cleaned up anyways, + # and we can start a new loop. + asyncio.run(main(args.host)) diff --git a/examples/async_reconnecting_ssl_client.py b/examples/async_reconnecting_ssl_client.py index 11d5817..c5f40d1 100755 --- a/examples/async_reconnecting_ssl_client.py +++ b/examples/async_reconnecting_ssl_client.py @@ -4,16 +4,14 @@ import asyncio import argparse import os import time -import socket import ssl +import socket import capnp import thread_capnp this_dir = os.path.dirname(os.path.abspath(__file__)) -capnp.remove_event_loop() -capnp.create_event_loop(threaded=True) def parse_args(): @@ -33,32 +31,10 @@ class StatusSubscriber(thread_capnp.Example.StatusSubscriber.Server): print("status: {}".format(time.time())) -async def myreader(client, reader): - while True: - try: - # Must be a wait_for in order to give watch_connection a slot - # to try again - data = await asyncio.wait_for(reader.read(4096), timeout=1.0) - except asyncio.TimeoutError: - continue - client.write(data) - - -async def mywriter(client, writer): - while True: - try: - # Must be a wait_for in order to give watch_connection a slot - # to try again - data = await asyncio.wait_for(client.read(4096), timeout=1.0) - writer.write(data.tobytes()) - except asyncio.TimeoutError: - continue - - async def watch_connection(cap): while True: try: - await asyncio.wait_for(cap.alive().a_wait(), timeout=5) + await asyncio.wait_for(cap.alive(), timeout=5) await asyncio.sleep(1) except asyncio.TimeoutError: print("Watch timeout!") @@ -68,14 +44,11 @@ async def watch_connection(cap): async def background(cap): subscriber = StatusSubscriber() - promise = cap.subscribeStatus(subscriber) - await promise.a_wait() + await cap.subscribeStatus(subscriber) async def main(host): - host = host.split(":") - addr = host[0] - port = host[1] + addr, port = host.split(":") # Setup SSL context ctx = ssl.create_default_context( @@ -85,46 +58,33 @@ async def main(host): # Handle both IPv4 and IPv6 cases try: print("Try IPv4") - reader, writer = await asyncio.open_connection( + stream = await capnp.AsyncIoStream.create_connection( addr, port, ssl=ctx, family=socket.AF_INET ) - except OSError: + except Exception: print("Try IPv6") - try: - reader, writer = await asyncio.open_connection( - addr, port, ssl=ctx, family=socket.AF_INET6 - ) - except OSError: - return False + stream = await capnp.AsyncIoStream.create_connection( + addr, port, ssl=ctx, family=socket.AF_INET6 + ) - # Start TwoPartyClient using TwoWayPipe (takes no arguments in this mode) - client = capnp.TwoPartyClient() + client = capnp.TwoPartyClient(stream) cap = client.bootstrap().cast_as(thread_capnp.Example) - # Start watcher to restart socket connection if it is lost - overalltasks = [] - watcher = [watch_connection(cap)] - overalltasks.append(asyncio.gather(*watcher, return_exceptions=True)) - - # Assemble reader and writer tasks, run in the background - coroutines = [myreader(client, reader), mywriter(client, writer)] - overalltasks.append(asyncio.gather(*coroutines, return_exceptions=True)) - - # Start background task for subscriber - tasks = [background(cap)] - overalltasks.append(asyncio.gather(*tasks, return_exceptions=True)) + # Start watcher to restart socket connection if it is lost and subscriber background task + background_tasks = asyncio.gather( + background(cap), watch_connection(cap), return_exceptions=True + ) # Run blocking tasks print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - for task in overalltasks: - task.cancel() + background_tasks.cancel() return True diff --git a/examples/async_server.py b/examples/async_server.py index 0a89b58..bacb3f6 100755 --- a/examples/async_server.py +++ b/examples/async_server.py @@ -3,7 +3,6 @@ import argparse import asyncio import logging -import socket import capnp import thread_capnp @@ -25,61 +24,12 @@ class ExampleImpl(thread_capnp.Example.Server): ) def longRunning(self, **kwargs): - return capnp.getTimer().after_delay(1 * 10**9) + return capnp.getTimer().after_delay(11 * 10**8) -class Server: - async def myreader(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.reader.read(4096), timeout=0.1) - except asyncio.TimeoutError: - logger.debug("myreader timeout.") - continue - except Exception as err: - logger.error("Unknown myreader err: %s", err) - return False - await self.server.write(data) - logger.debug("myreader done.") - return True - - async def mywriter(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.server.read(4096), timeout=0.1) - self.writer.write(data.tobytes()) - except asyncio.TimeoutError: - logger.debug("mywriter timeout.") - continue - except Exception as err: - logger.error("Unknown mywriter err: %s", err) - return False - logger.debug("mywriter done.") - return True - - async def myserver(self, reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - self.server = capnp.TwoPartyServer(bootstrap=ExampleImpl()) - self.reader = reader - self.writer = writer - self.retry = True - - # Assemble reader and writer tasks, run in the background - coroutines = [self.myreader(), self.mywriter()] - tasks = asyncio.gather(*coroutines, return_exceptions=True) - - while True: - self.server.poll_once() - # Check to see if reader has been sent an eof (disconnect) - if self.reader.at_eof(): - self.retry = False - break - await asyncio.sleep(0.01) - - # Make wait for reader/writer to finish (prevent possible resource leaks) - await tasks +async def new_connection(stream): + server = capnp.TwoPartyServer(stream, bootstrap=ExampleImpl()) + await server.on_disconnect() def parse_args(): @@ -93,29 +43,9 @@ given address/port ADDRESS. """ return parser.parse_args() -async def new_connection(reader, writer): - server = Server() - await server.myserver(reader, writer) - - async def main(): - address = parse_args().address - host = address.split(":") - addr = host[0] - port = host[1] - - # Handle both IPv4 and IPv6 cases - try: - print("Try IPv4") - server = await asyncio.start_server( - new_connection, addr, port, family=socket.AF_INET - ) - except Exception: - print("Try IPv6") - server = await asyncio.start_server( - new_connection, addr, port, family=socket.AF_INET6 - ) - + host, port = parse_args().address.split(":") + server = await capnp.AsyncIoStream.create_server(new_connection, host, port) async with server: await server.serve_forever() diff --git a/examples/async_ssl_calculator_client.py b/examples/async_ssl_calculator_client.py index b83d24c..342701a 100755 --- a/examples/async_ssl_calculator_client.py +++ b/examples/async_ssl_calculator_client.py @@ -3,8 +3,8 @@ import argparse import asyncio import os -import socket import ssl +import socket import capnp import calculator_capnp @@ -29,19 +29,6 @@ class PowerFunction(calculator_capnp.Calculator.Function.Server): return pow(params[0], params[1]) -async def myreader(client, reader): - while True: - data = await reader.read(4096) - client.write(data) - - -async def mywriter(client, writer): - while True: - data = await client.read(4096) - writer.write(data.tobytes()) - await writer.drain() - - def parse_args(): parser = argparse.ArgumentParser( usage="Connects to the Calculator server \ @@ -53,9 +40,7 @@ at the given address and does some RPCs" async def main(host): - host = host.split(":") - addr = host[0] - port = host[1] + addr, port = host.split(":") # Setup SSL context ctx = ssl.create_default_context( @@ -65,21 +50,16 @@ async def main(host): # Handle both IPv4 and IPv6 cases try: print("Try IPv4") - reader, writer = await asyncio.open_connection( + stream = await capnp.AsyncIoStream.create_connection( addr, port, ssl=ctx, family=socket.AF_INET ) except Exception: print("Try IPv6") - reader, writer = await asyncio.open_connection( + stream = await capnp.AsyncIoStream.create_connection( addr, port, ssl=ctx, family=socket.AF_INET6 ) - # Start TwoPartyClient using TwoWayPipe (takes no arguments in this mode) - client = capnp.TwoPartyClient() - - # Assemble reader and writer tasks, run in the background - coroutines = [myreader(client, reader), mywriter(client, writer)] - asyncio.gather(*coroutines, return_exceptions=True) + client = capnp.TwoPartyClient(stream) # Bootstrap the Calculator interface calculator = client.bootstrap().cast_as(calculator_capnp.Calculator) @@ -117,7 +97,7 @@ async def main(host): # Now that we've sent all the requests, wait for the response. Until this # point, we haven't waited at all! - response = await read_promise.a_wait() + response = await read_promise assert response.value == 123 print("PASS") @@ -155,7 +135,7 @@ async def main(host): eval_promise = request.send() read_promise = eval_promise.value.read() - response = await read_promise.a_wait() + response = await read_promise assert response.value == 101 print("PASS") @@ -219,8 +199,8 @@ async def main(host): add_5_promise = add_5_request.send().value.read() # Now wait for the results. - assert (await add_3_promise.a_wait()).value == 27 - assert (await add_5_promise.a_wait()).value == 29 + assert (await add_3_promise).value == 27 + assert (await add_5_promise).value == 29 print("PASS") @@ -301,8 +281,8 @@ async def main(host): g_eval_promise = g_eval_request.send().value.read() # Wait for the results. - assert (await f_eval_promise.a_wait()).value == 1234 - assert (await g_eval_promise.a_wait()).value == 4244 + assert (await f_eval_promise).value == 1234 + assert (await g_eval_promise).value == 4244 print("PASS") @@ -340,7 +320,7 @@ async def main(host): add_params[1].literal = 5 # Send the request and wait. - response = await request.send().value.read().a_wait() + response = await request.send().value.read() assert response.value == 512 print("PASS") diff --git a/examples/async_ssl_calculator_server.py b/examples/async_ssl_calculator_server.py index c92fef2..4404889 100755 --- a/examples/async_ssl_calculator_server.py +++ b/examples/async_ssl_calculator_server.py @@ -4,8 +4,8 @@ import argparse import asyncio import logging import os -import socket import ssl +import socket import capnp import calculator_capnp @@ -17,60 +17,6 @@ logger.setLevel(logging.DEBUG) this_dir = os.path.dirname(os.path.abspath(__file__)) -class Server: - async def myreader(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.reader.read(4096), timeout=0.1) - except asyncio.TimeoutError: - logger.debug("myreader timeout.") - continue - except Exception as err: - logger.error("Unknown myreader err: %s", err) - return False - await self.server.write(data) - logger.debug("myreader done.") - return True - - async def mywriter(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.server.read(4096), timeout=0.1) - self.writer.write(data.tobytes()) - except asyncio.TimeoutError: - logger.debug("mywriter timeout.") - continue - except Exception as err: - logger.error("Unknown mywriter err: %s", err) - return False - logger.debug("mywriter done.") - return True - - async def myserver(self, reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - self.server = capnp.TwoPartyServer(bootstrap=CalculatorImpl()) - self.reader = reader - self.writer = writer - self.retry = True - - # Assemble reader and writer tasks, run in the background - coroutines = [self.myreader(), self.mywriter()] - tasks = asyncio.gather(*coroutines, return_exceptions=True) - - while True: - self.server.poll_once() - # Check to see if reader has been sent an eof (disconnect) - if self.reader.at_eof(): - self.retry = False - break - await asyncio.sleep(0.01) - - # Make wait for reader/writer to finish (prevent possible resource leaks) - await tasks - - def read_value(value): """Helper function to asynchronously call read() on a Calculator::Value and return a promise for the result. (In the future, the generated code might @@ -195,16 +141,13 @@ given address/port ADDRESS. """ return parser.parse_args() -async def new_connection(reader, writer): - server = Server() - await server.myserver(reader, writer) +async def new_connection(stream): + server = capnp.TwoPartyServer(stream, bootstrap=CalculatorImpl()) + await server.on_disconnect() async def main(): - address = parse_args().address - host = address.split(":") - addr = host[0] - port = host[1] + host, port = parse_args().address.split(":") # Setup SSL context ctx = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH) @@ -216,13 +159,13 @@ async def main(): # Handle both IPv4 and IPv6 cases try: print("Try IPv4") - server = await asyncio.start_server( - new_connection, addr, port, ssl=ctx, family=socket.AF_INET + server = await capnp.AsyncIoStream.create_server( + new_connection, host, port, ssl=ctx, family=socket.AF_INET ) except Exception: print("Try IPv6") - server = await asyncio.start_server( - new_connection, addr, port, ssl=ctx, family=socket.AF_INET6 + server = await capnp.AsyncIoStream.create_server( + new_connection, host, port, ssl=ctx, family=socket.AF_INET6 ) async with server: diff --git a/examples/async_ssl_client.py b/examples/async_ssl_client.py index 8a7be8a..fa46062 100755 --- a/examples/async_ssl_client.py +++ b/examples/async_ssl_client.py @@ -3,9 +3,9 @@ import argparse import asyncio import os -import socket import ssl import time +import socket import capnp import thread_capnp @@ -30,29 +30,13 @@ class StatusSubscriber(thread_capnp.Example.StatusSubscriber.Server): print("status: {}".format(time.time())) -async def myreader(client, reader): - while True: - data = await reader.read(4096) - client.write(data) - - -async def mywriter(client, writer): - while True: - data = await client.read(4096) - writer.write(data.tobytes()) - await writer.drain() - - async def background(cap): subscriber = StatusSubscriber() - promise = cap.subscribeStatus(subscriber) - await promise.a_wait() + await cap.subscribeStatus(subscriber) async def main(host): - host = host.split(":") - addr = host[0] - port = host[1] + addr, port = host.split(":") # Setup SSL context ctx = ssl.create_default_context( @@ -62,34 +46,28 @@ async def main(host): # Handle both IPv4 and IPv6 cases try: print("Try IPv4") - reader, writer = await asyncio.open_connection( + stream = await capnp.AsyncIoStream.create_connection( addr, port, ssl=ctx, family=socket.AF_INET ) except Exception: print("Try IPv6") - reader, writer = await asyncio.open_connection( + stream = await capnp.AsyncIoStream.create_connection( addr, port, ssl=ctx, family=socket.AF_INET6 ) - # Start TwoPartyClient using TwoWayPipe (takes no arguments in this mode) - client = capnp.TwoPartyClient() + client = capnp.TwoPartyClient(stream) cap = client.bootstrap().cast_as(thread_capnp.Example) - # Assemble reader and writer tasks, run in the background - coroutines = [myreader(client, reader), mywriter(client, writer)] - asyncio.gather(*coroutines, return_exceptions=True) - # Start background task for subscriber - tasks = [background(cap)] - asyncio.gather(*tasks, return_exceptions=True) + asyncio.create_task(background(cap)) # Run blocking tasks print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) - await cap.longRunning().a_wait() + await cap.longRunning() print("main: {}".format(time.time())) diff --git a/examples/async_ssl_server.py b/examples/async_ssl_server.py index bb59e3c..fc3941a 100755 --- a/examples/async_ssl_server.py +++ b/examples/async_ssl_server.py @@ -4,8 +4,8 @@ import argparse import asyncio import logging import os -import socket import ssl +import socket import capnp import thread_capnp @@ -35,81 +35,21 @@ class ExampleImpl(thread_capnp.Example.Server): return True -class Server: - async def myreader(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.reader.read(4096), timeout=0.1) - except asyncio.TimeoutError: - logger.debug("myreader timeout.") - continue - except Exception as err: - logger.error("Unknown myreader err: %s", err) - return False - await self.server.write(data) - logger.debug("myreader done.") - return True - - async def mywriter(self): - while self.retry: - try: - # Must be a wait_for so we don't block on read() - data = await asyncio.wait_for(self.server.read(4096), timeout=0.1) - self.writer.write(data.tobytes()) - except asyncio.TimeoutError: - logger.debug("mywriter timeout.") - continue - except Exception as err: - logger.error("Unknown mywriter err: %s", err) - return False - logger.debug("mywriter done.") - return True - - async def myserver(self, reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - self.server = capnp.TwoPartyServer(bootstrap=ExampleImpl()) - self.reader = reader - self.writer = writer - self.retry = True - - # Assemble reader and writer tasks, run in the background - coroutines = [self.myreader(), self.mywriter()] - tasks = asyncio.gather(*coroutines, return_exceptions=True) - - while True: - self.server.poll_once() - # Check to see if reader has been sent an eof (disconnect) - if self.reader.at_eof(): - self.retry = False - break - await asyncio.sleep(0.01) - - # Make wait for reader/writer to finish (prevent possible resource leaks) - await tasks - - -async def new_connection(reader, writer): - server = Server() - await server.myserver(reader, writer) +async def new_connection(stream): + server = capnp.TwoPartyServer(stream, bootstrap=ExampleImpl()) + await server.on_disconnect() def parse_args(): parser = argparse.ArgumentParser( - usage="""Runs the server bound to the\ -given address/port ADDRESS. """ + usage="""Runs the server bound to the given address/port ADDRESS. """ ) - parser.add_argument("address", help="ADDRESS:PORT") - return parser.parse_args() async def main(): - address = parse_args().address - host = address.split(":") - addr = host[0] - port = host[1] + host, port = parse_args().address.split(":") # Setup SSL context ctx = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH) @@ -121,21 +61,13 @@ async def main(): # Handle both IPv4 and IPv6 cases try: print("Try IPv4") - server = await asyncio.start_server( - new_connection, - addr, - port, - ssl=ctx, - family=socket.AF_INET, + server = await capnp.AsyncIoStream.create_server( + new_connection, host, port, ssl=ctx, family=socket.AF_INET ) except Exception: print("Try IPv6") - server = await asyncio.start_server( - new_connection, - addr, - port, - ssl=ctx, - family=socket.AF_INET6, + server = await capnp.AsyncIoStream.create_server( + new_connection, host, port, ssl=ctx, family=socket.AF_INET6 ) async with server: diff --git a/examples/thread_client.py b/examples/thread_client.py index 19e89fd..317bc62 100755 --- a/examples/thread_client.py +++ b/examples/thread_client.py @@ -7,9 +7,6 @@ import capnp import thread_capnp -capnp.remove_event_loop() -capnp.create_event_loop(threaded=True) - def parse_args(): parser = argparse.ArgumentParser( diff --git a/examples/thread_server.py b/examples/thread_server.py index 6f102df..25b2ae0 100755 --- a/examples/thread_server.py +++ b/examples/thread_server.py @@ -2,7 +2,6 @@ import argparse import capnp -import time import thread_capnp @@ -38,9 +37,7 @@ def main(): address = parse_args().address server = capnp.TwoPartyServer(address, bootstrap=ExampleImpl()) - while True: - server.poll_once() - time.sleep(0.001) + server.run_forever() if __name__ == "__main__": diff --git a/setup.py b/setup.py index c79a51c..ef29917 100644 --- a/setup.py +++ b/setup.py @@ -200,7 +200,10 @@ import Cython # noqa: E402 extensions = [ Extension( "*", - ["capnp/helpers/capabilityHelper.cpp", "capnp/lib/*.pyx"], + [ + "capnp/helpers/capabilityHelper.cpp", + "capnp/lib/*.pyx", + ], extra_compile_args=extra_compile_args, extra_link_args=extra_link_args, language="c++", diff --git a/test/test_capability.py b/test/test_capability.py index bb06956..01c6f25 100644 --- a/test/test_capability.py +++ b/test/test_capability.py @@ -284,6 +284,21 @@ def test_cancel(): with pytest.raises(Exception): remote.wait() + req = client.foo(5) + trans = req.then(lambda x: 5) + req.cancel() # Cancel a promise that was already consumed + assert trans.wait() == 5 + + req = client.foo(5) + req.cancel() + with pytest.raises(Exception): + trans = req.then(lambda x: 5) + + req = client.foo(5) + assert req.wait().x == "26" + with pytest.raises(Exception): + req.wait() + def test_timer(): global test_timer_var @@ -350,6 +365,33 @@ def test_then_args(): client.foo(i=5).then(lambda x, y: 1) +class PromiseJoinServer(capability.TestPipeline.Server): + def getCap(self, n, inCap, _context, **kwargs): + def _then(response): + _results = _context.results + _results.s = response.x + "_bar" + _results.outBox.cap = inCap + + return ( + inCap.foo(i=n) + .then( + lambda res: capnp.Promise(int(res.x)) + ) # Make sure that Promise is flattened + .then( + lambda x: inCap.foo(i=x + 1) + ) # Make sure that RemotePromise is flattened + .then(_then) + ) + + +def test_promise_joining(): + client = capability.TestPipeline._new_client(PromiseJoinServer()) + foo_client = capability.TestInterface._new_client(Server()) + + remote = client.getCap(n=5, inCap=foo_client) + assert remote.wait().s == "136_bar" + + class ExtendsServer(Server): def qux(self, **kwargs): pass diff --git a/test/test_threads.py b/test/test_threads.py index 88f04fa..4756294 100644 --- a/test/test_threads.py +++ b/test/test_threads.py @@ -10,49 +10,9 @@ import pytest import capnp -from capnp.lib.capnp import KjException - import test_capability_capnp -@pytest.mark.skipif( - platform.python_implementation() == "PyPy", - reason="pycapnp's GIL handling isn't working properly at the moment for PyPy", -) -def test_making_event_loop(): - """ - Event loop test - """ - capnp.remove_event_loop(True) - capnp.create_event_loop() - - capnp.remove_event_loop() - capnp.create_event_loop() - - -@pytest.mark.skipif( - platform.python_implementation() == "PyPy", - reason="pycapnp's GIL handling isn't working properly at the moment for PyPy", -) -def test_making_threaded_event_loop(): - """ - Threaded event loop test - """ - # The following raises a KjException, and if not caught causes an SIGABRT: - # kj/async.c++:973: failed: expected head == nullptr; EventLoop destroyed with events still in the queue. - # Memory leak?; head->trace() = kj::_::ForkHub - # kj::_::AdapterPromiseNode > - # stack: ... - # python(..) malloc: *** error for object 0x...: pointer being freed was not allocated - # python(..) malloc: *** set a breakpoint in malloc_error_break to debug - # Fatal Python error: Aborted - capnp.remove_event_loop(KjException) - capnp.create_event_loop(KjException) - - capnp.remove_event_loop() - capnp.create_event_loop(KjException) - - class Server(test_capability_capnp.TestInterface.Server): """ Server @@ -76,9 +36,6 @@ def test_using_threads(): """ Thread test """ - capnp.remove_event_loop(True) - capnp.create_event_loop(True) - read, write = socket.socketpair() def run_server():