From c037342615c5f8336e67717963032b6de7960e58 Mon Sep 17 00:00:00 2001 From: Lasse Blaauwbroek Date: Wed, 19 Apr 2023 18:25:01 +0200 Subject: [PATCH] Allow capability implementation methods to be async --- capnp/helpers/capabilityHelper.cpp | 21 ++++++++++ capnp/helpers/capabilityHelper.h | 6 +++ capnp/helpers/helpers.pxd | 1 + capnp/includes/capnp_cpp.pxd | 1 + capnp/lib/capnp.pxd | 2 +- capnp/lib/capnp.pyx | 61 +++++++++++++++++++++-------- examples/async_calculator_server.py | 42 +++++++------------- examples/async_server.py | 17 ++++---- 8 files changed, 95 insertions(+), 56 deletions(-) diff --git a/capnp/helpers/capabilityHelper.cpp b/capnp/helpers/capabilityHelper.cpp index 92c0444..68f7753 100644 --- a/capnp/helpers/capabilityHelper.cpp +++ b/capnp/helpers/capabilityHelper.cpp @@ -203,6 +203,27 @@ void PyAsyncIoStream::shutdownWrite() { _asyncio_stream_shutdown_write(protocol->obj); } + +class TaskToPromiseAdapter { +public: + TaskToPromiseAdapter(kj::PromiseFulfiller& fulfiller, + kj::Own task, PyObject* callback) + : task(kj::mv(task)) { + promise_task_add_done_callback(this->task->obj, callback, fulfiller); + } + + ~TaskToPromiseAdapter() { + promise_task_cancel(this->task->obj); + } + +private: + kj::Own task; +}; + +kj::Promise taskToPromise(kj::Own task, PyObject* callback) { + return kj::newAdaptedPromise(kj::mv(task), callback); +} + void init_capnp_api() { import_capnp__lib__capnp(); } diff --git a/capnp/helpers/capabilityHelper.h b/capnp/helpers/capabilityHelper.h index d41bad9..76e8c50 100644 --- a/capnp/helpers/capabilityHelper.h +++ b/capnp/helpers/capabilityHelper.h @@ -141,4 +141,10 @@ inline void rejectVoidDisconnected(kj::PromiseFulfiller& fulfiller, kj::St fulfiller.reject(KJ_EXCEPTION(DISCONNECTED, message)); } +inline kj::Exception makeException(kj::StringPtr message) { + return KJ_EXCEPTION(FAILED, message); +} + +kj::Promise taskToPromise(kj::Own coroutine, PyObject* callback); + void init_capnp_api(); diff --git a/capnp/helpers/helpers.pxd b/capnp/helpers/helpers.pxd index 45e3f4c..5682d12 100644 --- a/capnp/helpers/helpers.pxd +++ b/capnp/helpers/helpers.pxd @@ -32,6 +32,7 @@ cdef extern from "capnp/helpers/capabilityHelper.h": PyPromise convert_to_pypromise(Own[VoidPromise]) VoidPromise convert_to_voidpromise(Own[PyPromise]) PyPromise wrapSizePromise(Promise[size_t]) + VoidPromise taskToPromise(Own[PyRefCounter] coroutine, PyObject* callback) void init_capnp_api() cdef extern from "capnp/helpers/rpcHelper.h": diff --git a/capnp/includes/capnp_cpp.pxd b/capnp/includes/capnp_cpp.pxd index 346533c..563680b 100644 --- a/capnp/includes/capnp_cpp.pxd +++ b/capnp/includes/capnp_cpp.pxd @@ -564,3 +564,4 @@ cdef extern from "capnp/helpers/capabilityHelper.h": PyAsyncIoStream(PyObject* thisptr) void rejectDisconnected[T](PromiseFulfiller[T]& fulfiller, StringPtr message) void rejectVoidDisconnected(VoidPromiseFulfiller& fulfiller, StringPtr message) + Exception makeException(StringPtr message) diff --git a/capnp/lib/capnp.pxd b/capnp/lib/capnp.pxd index 21a3e41..2e319d3 100644 --- a/capnp/lib/capnp.pxd +++ b/capnp/lib/capnp.pxd @@ -160,7 +160,7 @@ cdef _setDynamicFieldStatic(DynamicStruct_Builder thisptr, field, value, parent) cdef api object wrap_dynamic_struct_reader(Response & r) with gil cdef api Promise[void] * call_server_method( - PyObject * _server, char * _method_name, CallContext & _context) except * with gil + object server, char * _method_name, CallContext & _context) except * with gil cdef api convert_array_pyobject(PyArray & arr) 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 diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index c8c758d..b6df98f 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -10,7 +10,7 @@ cimport cython # noqa: E402 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 capnp.includes.capnp_cpp cimport AsyncIoStream, WaitScope, PyPromise, VoidPromise, EventPort, EventLoop, WaitScope, LowLevelAsyncIoProvider, AsyncIoProvider, newAsyncIoProvider, MonotonicClock, Timer, TimerImpl, systemPreciseMonotonicClock, MILLISECONDS, Canceler, PyAsyncIoStream, PromiseFulfiller, VoidPromiseFulfiller, makeException from cpython cimport array, Py_buffer, PyObject_CheckBuffer, memoryview, buffer from cpython.buffer cimport PyBUF_SIMPLE, PyBUF_WRITABLE @@ -68,10 +68,39 @@ cdef api object wrap_remote_call(object func, Response & r): cdef _find_field_order(struct_node): return [f.name for f in sorted(struct_node.fields, key=_attrgetter('codeOrder'))] +cdef class _VoidPromiseFulfiller: + cdef VoidPromiseFulfiller* fulfiller -cdef api VoidPromise * call_server_method(PyObject * _server, + cdef _init(self, VoidPromiseFulfiller* fulfiller): + self.fulfiller = fulfiller + return self + +def void_task_done_callback(method_name, _VoidPromiseFulfiller fulfiller, task): + if task.cancelled(): + fulfiller.fulfiller.reject(makeException(capnp.StringPtr( + f"Server task for method {method_name} was cancelled"))) + return + + exc = task.exception() + if exc is not None: + fulfiller.fulfiller.reject(makeException(capnp.StringPtr(str(exc)))) + return + + res = task.result() + if res is not None: + fulfiller.fulfiller.reject(makeException(capnp.StringPtr( + f"Async server function ({method_name}) returned a non-none value: return = {res}"))) + else: + fulfiller.fulfiller.fulfill() + +cdef api void promise_task_add_done_callback(object task, object callback, VoidPromiseFulfiller& fulfiller): + task.add_done_callback(_partial(callback, _VoidPromiseFulfiller()._init(&fulfiller))) + +cdef api void promise_task_cancel(object task): + task.cancel() + +cdef api VoidPromise * call_server_method(object server, char * _method_name, CallContext & _context) except * with gil: - server = _server method_name = _method_name context = _CallContext()._init(_context) # TODO:MEMORY: invalidate this with promise chain @@ -83,6 +112,12 @@ cdef api VoidPromise * call_server_method(PyObject * _server, return new VoidPromise(moveVoidPromise(deref((<_VoidPromise>ret).thisptr))) elif type(ret) is _Promise: return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) + elif asyncio.iscoroutine(ret): + task = asyncio.create_task(ret) + callback = _partial(void_task_done_callback, method_name) + return new VoidPromise(helpers.taskToPromise( + capnp.heap[PyRefCounter](task), + callback)) else: try: warning_msg = ( @@ -93,20 +128,6 @@ cdef api VoidPromise * call_server_method(PyObject * _server, _warnings.warn_explicit( warning_msg, UserWarning, _inspect.getsourcefile(func), _inspect.getsourcelines(func)[1]) - if ret is not None: - if type(ret) is _Promise: - return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) - elif type(ret) is _Promise: - return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) - else: - try: - warning_msg = ( - "Server function ({}) returned a value that was not a Promise: return = {}" - .format(method_name, str(ret))) - except Exception: - warning_msg = 'Server function (%s) returned a value that was not a Promise' % (method_name) - _warnings.warn_explicit( - warning_msg, UserWarning, _inspect.getsourcefile(func), _inspect.getsourcelines(func)[1]) else: func = getattr(server, method_name) # will raise if no function found params = context.params @@ -119,6 +140,12 @@ cdef api VoidPromise * call_server_method(PyObject * _server, return new VoidPromise(moveVoidPromise(deref((<_VoidPromise>ret).thisptr))) elif type(ret) is _Promise: return new VoidPromise(helpers.convert_to_voidpromise(move((<_Promise>ret).thisptr))) + elif asyncio.iscoroutine(ret): + task = asyncio.create_task(ret) + callback = _partial(void_task_done_callback, method_name) + return new VoidPromise(helpers.taskToPromise( + capnp.heap[PyRefCounter](task), + callback)) if not isinstance(ret, tuple): ret = (ret,) names = _find_field_order(context.results.schema.node.struct) diff --git a/examples/async_calculator_server.py b/examples/async_calculator_server.py index f1c6dc4..6c1b40b 100755 --- a/examples/async_calculator_server.py +++ b/examples/async_calculator_server.py @@ -12,15 +12,7 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.DEBUG) -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 - include something like this automatically.)""" - - return value.read().then(lambda result: result.value) - - -def evaluate_impl(expression, params=None): +async def evaluate_impl(expression, params=None): """Implementation of CalculatorImpl::evaluate(), also shared by FunctionImpl::call(). In the latter case, `params` are the parameter values passed to the function; in the former case, `params` is just an @@ -29,26 +21,23 @@ def evaluate_impl(expression, params=None): which = expression.which() if which == "literal": - return capnp.Promise(expression.literal) + return expression.literal elif which == "previousResult": - return read_value(expression.previousResult) + return (await expression.previousResult.read()).value elif which == "parameter": assert expression.parameter < len(params) - return capnp.Promise(params[expression.parameter]) + return params[expression.parameter] elif which == "call": call = expression.call func = call.function # Evaluate each parameter. paramPromises = [evaluate_impl(param, params) for param in call.params] + vals = await asyncio.gather(*paramPromises) - joinedParams = capnp.join_promises(paramPromises) # When the parameters are complete, call the function. - ret = joinedParams.then(lambda vals: func.call(vals)).then( - lambda result: result.value - ) - - return ret + result = await func.call(vals) + return result.value else: raise ValueError("Unknown expression type: " + which) @@ -72,17 +61,15 @@ class FunctionImpl(calculator_capnp.Calculator.Function.Server): self.paramCount = paramCount self.body = body.as_builder() - def call(self, params, _context, **kwargs): + async def call(self, params, _context, **kwargs): """Note that we're returning a Promise object here, and bypassing the helper functionality that normally sets the results struct from the returned object. Instead, we set _context.results directly inside of another promise""" assert len(params) == self.paramCount - # using setattr because '=' is not allowed inside of lambdas - return evaluate_impl(self.body, params).then( - lambda value: setattr(_context.results, "value", value) - ) + value = await evaluate_impl(self.body, params) + _context.results.value = value class OperatorImpl(calculator_capnp.Calculator.Function.Server): @@ -113,10 +100,9 @@ class OperatorImpl(calculator_capnp.Calculator.Function.Server): class CalculatorImpl(calculator_capnp.Calculator.Server): "Implementation of the Calculator Cap'n Proto interface." - def evaluate(self, expression, _context, **kwargs): - return evaluate_impl(expression).then( - lambda value: setattr(_context.results, "value", ValueImpl(value)) - ) + async def evaluate(self, expression, _context, **kwargs): + value = await evaluate_impl(expression) + _context.results.value = ValueImpl(value) def defFunction(self, paramCount, body, _context, **kwargs): return FunctionImpl(paramCount, body) @@ -133,7 +119,7 @@ async def new_connection(stream): def parse_args(): parser = argparse.ArgumentParser( usage="""Runs the server bound to the\ -given address/port ADDRESS. """ + given address/port ADDRESS. """ ) parser.add_argument("address", help="ADDRESS:PORT") diff --git a/examples/async_server.py b/examples/async_server.py index bacb3f6..00e85f6 100755 --- a/examples/async_server.py +++ b/examples/async_server.py @@ -15,16 +15,13 @@ logger.setLevel(logging.DEBUG) class ExampleImpl(thread_capnp.Example.Server): "Implementation of the Example threading Cap'n Proto interface." - def subscribeStatus(self, subscriber, **kwargs): - return ( - capnp.getTimer() - .after_delay(10**9) - .then(lambda: subscriber.status(True)) - .then(lambda _: self.subscribeStatus(subscriber)) - ) + async def subscribeStatus(self, subscriber, **kwargs): + await asyncio.sleep(1) + await subscriber.status(True) + await self.subscribeStatus(subscriber) - def longRunning(self, **kwargs): - return capnp.getTimer().after_delay(11 * 10**8) + async def longRunning(self, **kwargs): + await asyncio.sleep(1) async def new_connection(stream): @@ -35,7 +32,7 @@ async def new_connection(stream): def parse_args(): parser = argparse.ArgumentParser( usage="""Runs the server bound to the\ -given address/port ADDRESS. """ + given address/port ADDRESS. """ ) parser.add_argument("address", help="ADDRESS:PORT")