From d53aa24733584ce864a08fe1c6b668ac99c15c0e Mon Sep 17 00:00:00 2001 From: Lasse Blaauwbroek Date: Wed, 19 Apr 2023 18:49:19 +0200 Subject: [PATCH] Allow async capability implementation methods to return a tuple --- capnp/lib/capnp.pyx | 32 ++++++++++++++++++----------- examples/async_calculator_server.py | 6 ++---- 2 files changed, 22 insertions(+), 16 deletions(-) diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index b6df98f..8401d86 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -99,6 +99,21 @@ cdef api void promise_task_add_done_callback(object task, object callback, VoidP cdef api void promise_task_cancel(object task): task.cancel() +def fill_context(method_name, context, returned_data): + if returned_data is None: + return + if not isinstance(returned_data, tuple): + returned_data = (returned_data,) + names = _find_field_order(context.results.schema.node.struct) + if len(returned_data) > len(names): + raise KjException( + "Too many values returned from `{}`. Expected {} and got {}" + .format(method_name, len(names), len(returned_data))) + + results = context.results + for arg_name, arg_val in zip(names, returned_data): + setattr(results, arg_name, arg_val) + cdef api VoidPromise * call_server_method(object server, char * _method_name, CallContext & _context) except * with gil: method_name = _method_name @@ -141,22 +156,15 @@ cdef api VoidPromise * call_server_method(object server, 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) + async def finalize(): + fill_context(method_name, context, await ret) + task = asyncio.create_task(finalize()) 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) - if len(ret) > len(names): - raise KjException( - "Too many values returned from `{}`. Expected {} and got {}" - .format(method_name, len(names), len(ret))) - - results = context.results - for arg_name, arg_val in zip(names, ret): - setattr(results, arg_name, arg_val) + else: + fill_context(method_name, context, ret) return NULL diff --git a/examples/async_calculator_server.py b/examples/async_calculator_server.py index 6c1b40b..3e00a00 100755 --- a/examples/async_calculator_server.py +++ b/examples/async_calculator_server.py @@ -68,8 +68,7 @@ class FunctionImpl(calculator_capnp.Calculator.Function.Server): another promise""" assert len(params) == self.paramCount - value = await evaluate_impl(self.body, params) - _context.results.value = value + return await evaluate_impl(self.body, params) class OperatorImpl(calculator_capnp.Calculator.Function.Server): @@ -101,8 +100,7 @@ class CalculatorImpl(calculator_capnp.Calculator.Server): "Implementation of the Calculator Cap'n Proto interface." async def evaluate(self, expression, _context, **kwargs): - value = await evaluate_impl(expression) - _context.results.value = ValueImpl(value) + return ValueImpl(await evaluate_impl(expression)) def defFunction(self, paramCount, body, _context, **kwargs): return FunctionImpl(paramCount, body)