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
This commit is contained in:
Jacob Alexander
2020-01-08 18:06:26 -08:00
parent ea0858352b
commit dd3551f978
4 changed files with 211 additions and 69 deletions

View File

@@ -2406,6 +2406,9 @@ cdef class TwoPartyServer:
cpdef on_disconnect(self) except +reraise_kj_exception: cpdef on_disconnect(self) except +reraise_kj_exception:
return _VoidPromise()._init(deref(self._network.thisptr).onDisconnect()) return _VoidPromise()._init(deref(self._network.thisptr).onDisconnect())
def poll_once(self):
return poll_once()
async def poll_forever(self): async def poll_forever(self):
while True: while True:
poll_once() poll_once()

View File

@@ -1,25 +1,80 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
from __future__ import print_function from __future__ import print_function
import argparse import argparse
import asyncio import asyncio
import logging
import socket import socket
import capnp
import capnp
import calculator_capnp import calculator_capnp
async def myreader(client, reader): logger = logging.getLogger(__name__)
while True: logger.setLevel(logging.DEBUG)
data = await reader.read(4096)
await client.write(data)
async def mywriter(client, writer): 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: while True:
data = await client.read(4096) self.server.poll_once()
writer.write(data.tobytes()) # Check to see if reader has been sent an eof (disconnect)
await writer.drain() 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): def read_value(value):
@@ -142,15 +197,9 @@ given address/port ADDRESS. ''')
return parser.parse_args() return parser.parse_args()
async def myserver(reader, writer): async def new_connection(reader, writer):
# Start TwoPartyServer using TwoWayPipe (only requires bootstrap) server = Server()
server = capnp.TwoPartyServer(bootstrap=CalculatorImpl()) await server.myserver(reader, writer)
# 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 main(): async def main():
@@ -163,13 +212,13 @@ async def main():
try: try:
print("Try IPv4") print("Try IPv4")
server = await asyncio.start_server( server = await asyncio.start_server(
myserver, new_connection,
addr, port, addr, port,
) )
except Exception: except Exception:
print("Try IPv6") print("Try IPv6")
server = await asyncio.start_server( server = await asyncio.start_server(
myserver, new_connection,
addr, port, addr, port,
family=socket.AF_INET6 family=socket.AF_INET6
) )

View File

@@ -3,12 +3,17 @@
from __future__ import print_function from __future__ import print_function
import argparse import argparse
import capnp
import thread_capnp
import asyncio import asyncio
import logging
import socket import socket
import capnp
import thread_capnp
logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
class ExampleImpl(thread_capnp.Example.Server): class ExampleImpl(thread_capnp.Example.Server):
@@ -23,30 +28,66 @@ class ExampleImpl(thread_capnp.Example.Server):
return capnp.getTimer().after_delay(1 * 10**9) return capnp.getTimer().after_delay(1 * 10**9)
async def myreader(server, reader): class Server:
while True: async def myreader(self):
data = await reader.read(4096) while self.retry:
# Close connection if 0 bytes read try:
if len(data) == 0: # Must be a wait_for so we don't block on read()
server.close() data = await asyncio.wait_for(
await server.write(data) 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): async def mywriter(self):
while True: while self.retry:
data = await server.read(4096) try:
writer.write(data.tobytes()) # 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): async def myserver(self, reader, writer):
# Start TwoPartyServer using TwoWayPipe (only requires bootstrap) # Start TwoPartyServer using TwoWayPipe (only requires bootstrap)
server = capnp.TwoPartyServer(bootstrap=ExampleImpl()) self.server = capnp.TwoPartyServer(bootstrap=ExampleImpl())
self.reader = reader
self.writer = writer
self.retry = True
# Assemble reader and writer tasks, run in the background # Assemble reader and writer tasks, run in the background
coroutines = [myreader(server, reader), mywriter(server, writer)] coroutines = [self.myreader(), self.mywriter()]
asyncio.gather(*coroutines, return_exceptions=True) 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(): def parse_args():
@@ -58,6 +99,11 @@ given address/port ADDRESS. ''')
return parser.parse_args() return parser.parse_args()
async def new_connection(reader, writer):
server = Server()
await server.myserver(reader, writer)
async def main(): async def main():
address = parse_args().address address = parse_args().address
host = address.split(':') host = address.split(':')
@@ -68,13 +114,13 @@ async def main():
try: try:
print("Try IPv4") print("Try IPv4")
server = await asyncio.start_server( server = await asyncio.start_server(
myserver, new_connection,
addr, port, addr, port,
) )
except Exception: except Exception:
print("Try IPv6") print("Try IPv6")
server = await asyncio.start_server( server = await asyncio.start_server(
myserver, new_connection,
addr, port, addr, port,
family=socket.AF_INET6 family=socket.AF_INET6
) )

View File

@@ -3,14 +3,18 @@
from __future__ import print_function from __future__ import print_function
import argparse import argparse
import os
import capnp
import thread_capnp
import asyncio import asyncio
import logging
import os
import socket import socket
import ssl import ssl
import capnp
import thread_capnp
logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
this_dir = os.path.dirname(os.path.abspath(__file__)) this_dir = os.path.dirname(os.path.abspath(__file__))
@@ -31,31 +35,71 @@ class ExampleImpl(thread_capnp.Example.Server):
return True return True
async def myreader(server, reader): class Server:
while True: async def myreader(self):
data = await reader.read(4096) while self.retry:
# Close connection if 0 bytes read try:
if len(data) == 0: # Must be a wait_for so we don't block on read()
server.close() data = await asyncio.wait_for(
await server.write(data) 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): async def mywriter(self):
while True: while self.retry:
data = await server.read(4096) try:
writer.write(data.tobytes()) # Must be a wait_for so we don't block on read()
await writer.drain() 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): async def myserver(self, reader, writer):
# Start TwoPartyServer using TwoWayPipe (only requires bootstrap) # Start TwoPartyServer using TwoWayPipe (only requires bootstrap)
server = capnp.TwoPartyServer(bootstrap=ExampleImpl()) self.server = capnp.TwoPartyServer(bootstrap=ExampleImpl())
self.reader = reader
self.writer = writer
self.retry = True
# Assemble reader and writer tasks, run in the background # Assemble reader and writer tasks, run in the background
coroutines = [myreader(server, reader), mywriter(server, writer)] coroutines = [self.myreader(), self.mywriter()]
asyncio.gather(*coroutines, return_exceptions=True) 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(): def parse_args():
@@ -81,14 +125,14 @@ async def main():
try: try:
print("Try IPv4") print("Try IPv4")
server = await asyncio.start_server( server = await asyncio.start_server(
myserver, new_connection,
addr, port, addr, port,
ssl=ctx, ssl=ctx,
) )
except Exception: except Exception:
print("Try IPv6") print("Try IPv6")
server = await asyncio.start_server( server = await asyncio.start_server(
myserver, new_connection,
addr, port, addr, port,
ssl=ctx, ssl=ctx,
family=socket.AF_INET6, family=socket.AF_INET6,