From dd3551f9782ccba1ec3db22733dff266a6d692e8 Mon Sep 17 00:00:00 2001 From: Jacob Alexander Date: Wed, 8 Jan 2020 18:06:26 -0800 Subject: [PATCH] Adding poll_once() to TwoPartyServer API - poll_forever() doesn't allow for checking socket connection to client * Need to check for client, if eof received then we can close the connection and cleanup pycapnp cpu resources (async tasks) - Updated async examples to fix bugs * Add checking code for socket connection to the client and gracefully cleanup resources once the client socket connection closes * Instanciate a new TwoPartyServer per connection (allows for multiple connections) - Should resolve Issue #198 --- capnp/lib/capnp.pyx | 3 + examples/async_calculator_server.py | 91 +++++++++++++++++++++------- examples/async_server.py | 92 +++++++++++++++++++++------- examples/async_ssl_server.py | 94 +++++++++++++++++++++-------- 4 files changed, 211 insertions(+), 69 deletions(-) diff --git a/capnp/lib/capnp.pyx b/capnp/lib/capnp.pyx index dfcab60..414f19a 100644 --- a/capnp/lib/capnp.pyx +++ b/capnp/lib/capnp.pyx @@ -2406,6 +2406,9 @@ cdef class TwoPartyServer: cpdef on_disconnect(self) except +reraise_kj_exception: return _VoidPromise()._init(deref(self._network.thisptr).onDisconnect()) + def poll_once(self): + return poll_once() + async def poll_forever(self): while True: poll_once() diff --git a/examples/async_calculator_server.py b/examples/async_calculator_server.py index 8324902..0ecd944 100755 --- a/examples/async_calculator_server.py +++ b/examples/async_calculator_server.py @@ -1,25 +1,80 @@ #!/usr/bin/env python3 from __future__ import print_function + import argparse import asyncio +import logging import socket -import capnp +import capnp import calculator_capnp -async def myreader(client, reader): - while True: - data = await reader.read(4096) - await client.write(data) +logger = logging.getLogger(__name__) +logger.setLevel(logging.DEBUG) -async def mywriter(client, writer): - while True: - data = await client.read(4096) - writer.write(data.tobytes()) - await writer.drain() +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): @@ -142,15 +197,9 @@ given address/port ADDRESS. ''') return parser.parse_args() -async def myserver(reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - server = capnp.TwoPartyServer(bootstrap=CalculatorImpl()) - - # Assemble reader and writer tasks, run in the background - coroutines = [myreader(server, reader), mywriter(server, writer)] - asyncio.gather(*coroutines, return_exceptions=True) - - await server.poll_forever() +async def new_connection(reader, writer): + server = Server() + await server.myserver(reader, writer) async def main(): @@ -163,13 +212,13 @@ async def main(): try: print("Try IPv4") server = await asyncio.start_server( - myserver, + new_connection, addr, port, ) except Exception: print("Try IPv6") server = await asyncio.start_server( - myserver, + new_connection, addr, port, family=socket.AF_INET6 ) diff --git a/examples/async_server.py b/examples/async_server.py index e6fa225..acd4e1d 100755 --- a/examples/async_server.py +++ b/examples/async_server.py @@ -3,12 +3,17 @@ from __future__ import print_function import argparse -import capnp - -import thread_capnp import asyncio +import logging import socket +import capnp +import thread_capnp + + +logger = logging.getLogger(__name__) +logger.setLevel(logging.DEBUG) + class ExampleImpl(thread_capnp.Example.Server): @@ -23,30 +28,66 @@ class ExampleImpl(thread_capnp.Example.Server): return capnp.getTimer().after_delay(1 * 10**9) -async def myreader(server, reader): - while True: - data = await reader.read(4096) - # Close connection if 0 bytes read - if len(data) == 0: - server.close() - await server.write(data) +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(server, writer): - while True: - data = await server.read(4096) - writer.write(data.tobytes()) + 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(reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - server = capnp.TwoPartyServer(bootstrap=ExampleImpl()) + 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 = [myreader(server, reader), mywriter(server, writer)] - asyncio.gather(*coroutines, return_exceptions=True) + # Assemble reader and writer tasks, run in the background + coroutines = [self.myreader(), self.mywriter()] + tasks = asyncio.gather(*coroutines, return_exceptions=True) - await server.poll_forever() + 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 parse_args(): @@ -58,6 +99,11 @@ 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(':') @@ -68,13 +114,13 @@ async def main(): try: print("Try IPv4") server = await asyncio.start_server( - myserver, + new_connection, addr, port, ) except Exception: print("Try IPv6") server = await asyncio.start_server( - myserver, + new_connection, addr, port, family=socket.AF_INET6 ) diff --git a/examples/async_ssl_server.py b/examples/async_ssl_server.py index 06a7d28..2b8a958 100755 --- a/examples/async_ssl_server.py +++ b/examples/async_ssl_server.py @@ -3,14 +3,18 @@ from __future__ import print_function import argparse -import os -import capnp - -import thread_capnp import asyncio +import logging +import os import socket import ssl +import capnp +import thread_capnp + + +logger = logging.getLogger(__name__) +logger.setLevel(logging.DEBUG) this_dir = os.path.dirname(os.path.abspath(__file__)) @@ -31,31 +35,71 @@ class ExampleImpl(thread_capnp.Example.Server): return True -async def myreader(server, reader): - while True: - data = await reader.read(4096) - # Close connection if 0 bytes read - if len(data) == 0: - server.close() - await server.write(data) +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(server, writer): - while True: - data = await server.read(4096) - writer.write(data.tobytes()) - await writer.drain() + 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(reader, writer): - # Start TwoPartyServer using TwoWayPipe (only requires bootstrap) - server = capnp.TwoPartyServer(bootstrap=ExampleImpl()) + 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 = [myreader(server, reader), mywriter(server, writer)] - asyncio.gather(*coroutines, return_exceptions=True) + # Assemble reader and writer tasks, run in the background + coroutines = [self.myreader(), self.mywriter()] + tasks = asyncio.gather(*coroutines, return_exceptions=True) - await server.poll_forever() + 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) def parse_args(): @@ -81,14 +125,14 @@ async def main(): try: print("Try IPv4") server = await asyncio.start_server( - myserver, + new_connection, addr, port, ssl=ctx, ) except Exception: print("Try IPv6") server = await asyncio.start_server( - myserver, + new_connection, addr, port, ssl=ctx, family=socket.AF_INET6,