diff --git a/Doc/library/asyncio-stream.rst b/Doc/library/asyncio-stream.rst index e686a6a1c4cd32..717be22cfb3364 100644 --- a/Doc/library/asyncio-stream.rst +++ b/Doc/library/asyncio-stream.rst @@ -313,6 +313,28 @@ StreamWriter .. versionadded:: 3.7 + .. coroutinemethod:: sendfile(file, offset=0, count=None, \*, \ + fallback=True) + + Send a *file* to a peer. Return the total number of bytes + sent. + + For more details about arguments and implementation see + :meth:`loop.sendfile`. + + .. versionadded:: 3.8 + + .. coroutinemethod:: start_tls(sslcontext, \*, \ + server_hostname=None, \ + ssl_handshake_timeout=None) + + Upgrade an existing transport-based connection to TLS. + + For more details about arguments and implementation see + :meth:`loop.start_tls`. + + .. versionadded:: 3.8 + Examples ======== diff --git a/Lib/asyncio/streams.py b/Lib/asyncio/streams.py index 0afc66a473d418..8817828ee57a8e 100644 --- a/Lib/asyncio/streams.py +++ b/Lib/asyncio/streams.py @@ -1,13 +1,15 @@ __all__ = ( 'StreamReader', 'StreamWriter', 'StreamReaderProtocol', - 'open_connection', 'start_server') + 'open_connection', 'start_server', + 'connect') import socket import sys import weakref if hasattr(socket, 'AF_UNIX'): - __all__ += ('open_unix_connection', 'start_unix_server') + __all__ += ('open_unix_connection', 'start_unix_server', + 'connect_unix') from . import coroutines from . import events @@ -21,6 +23,23 @@ _DEFAULT_LIMIT = 2 ** 16 # 64 KiB +async def connect(host=None, port=None, *, + limit=_DEFAULT_LIMIT, **kwds): + assert 'loop' not in kwds + loop = events.get_running_loop() + stream = Stream(limit=limit, loop=loop) + protocol = StreamReaderProtocol(stream, loop=loop) + stream._set_protocol(protocol) + transport, _ = await loop.create_connection( + lambda: protocol, host, port, **kwds) + return stream + + +async def serve(client_connected_cb, host=None, port=None, *, + loop=None, limit=_DEFAULT_LIMIT, **kwds): + pass + + async def open_connection(host=None, port=None, *, loop=None, limit=_DEFAULT_LIMIT, **kwds): """A wrapper for create_connection() returning a (reader, writer) pair. @@ -88,6 +107,17 @@ def factory(): if hasattr(socket, 'AF_UNIX'): # UNIX Domain Sockets are supported on this platform + async def connect_unix(path=None, *, + limit=_DEFAULT_LIMIT, **kwds): + assert 'loop' not in kwds + loop = events.get_running_loop() + stream = Stream(limit=limit, loop=loop) + protocol = StreamReaderProtocol(stream, loop=loop) + stream._set_protocol(protocol) + transport, _ = await loop.create_unix_connection( + lambda: protocol, path, **kwds) + return stream + async def open_unix_connection(path=None, *, loop=None, limit=_DEFAULT_LIMIT, **kwds): """Similar to `open_connection` but works with UNIX Domain Sockets.""" @@ -417,7 +447,7 @@ def __init__(self, limit=_DEFAULT_LIMIT, loop=None): sys._getframe(1)) def __repr__(self): - info = ['StreamReader'] + info = [self.__class__.__name__] if self._buffer: info.append(f'{len(self._buffer)} bytes') if self._eof: @@ -742,3 +772,41 @@ async def __anext__(self): if val == b'': raise StopAsyncIteration return val + + +class Stream(StreamReader, StreamWriter): + def __init__(self, limit, loop): + StreamReader.__init__(self, limit, loop) + # Emulate StreamWriter ctor without an actual call + self._protocol = None # setup the attribute in _set_protocol() + + @property + def _reader(self): + # A trick for making StreamWriter work: the class requires + # self._reader attribute + return self + + def _set_protocol(self, protocol): + # a post-init method to set protocol instance + self._protocol = protocol + + def __repr__(self): + return StreamReader.__repr__(self) + + async def sendfile(self, file, offset=0, count=None, *, fallback=True): + await self.drain() + return await self._loop.sendfile(self._transport, file, + offset, count, fallback=fallback) + + async def start_tls(self, sslcontext, *, + server_hostname=None, + ssl_handshake_timeout=None): + server_side = self._protocol._client_connected_cb is not None + await self.drain() + transport = await self._loop.start_tls( + self._transport, self._protocol, sslcontext, + server_side=server_side, server_hostname=server_hostname, + ssl_handshake_timeout=ssl_handshake_timeout) + self._transport = transport + self._protocol._transport = transport + self._over_ssl = True diff --git a/Lib/test/test_asyncio/test_streams.py b/Lib/test/test_asyncio/test_streams.py index 0141df729ce080..ab2dda1c280f0a 100644 --- a/Lib/test/test_asyncio/test_streams.py +++ b/Lib/test/test_asyncio/test_streams.py @@ -987,6 +987,87 @@ def test_async_writer_api(self): self.assertEqual(messages, []) + def _basetest_connect(self, stream): + messages = [] + self.loop.set_exception_handler(lambda loop, ctx: messages.append(ctx)) + + stream.write(b'GET / HTTP/1.0\r\n\r\n') + f = stream.readline() + data = self.loop.run_until_complete(f) + self.assertEqual(data, b'HTTP/1.0 200 OK\r\n') + f = stream.read() + data = self.loop.run_until_complete(f) + self.assertTrue(data.endswith(b'\r\n\r\nTest message')) + stream.close() + self.loop.run_until_complete(stream.wait_closed()) + + self.assertEqual([], messages) + + def test_connect(self): + with test_utils.run_test_server() as httpd: + stream = self.loop.run_until_complete( + asyncio.connect(*httpd.address)) + self._basetest_connect(stream) + + @support.skip_unless_bind_unix_socket + def test_connect_unix(self): + with test_utils.run_test_unix_server() as httpd: + stream = self.loop.run_until_complete( + asyncio.connect_unix(httpd.address)) + self._basetest_connect(stream) + + def test_sendfile(self): + messages = [] + self.loop.set_exception_handler(lambda loop, ctx: messages.append(ctx)) + + with open(support.TESTFN, 'wb') as fp: + fp.write(b'data\n') + self.addCleanup(support.unlink, support.TESTFN) + + async def do_serve(reader, writer): + data = await reader.readline() + self.assertEqual(data, b'begin\n') + data = await reader.readline() + self.assertEqual(data, b'data\n') + data = await reader.readline() + self.assertEqual(data, b'end\n') + await writer.awrite(b'done\n') + await writer.aclose() + + server = self.loop.run_until_complete( + asyncio.start_server(do_serve, 'localhost', 0, loop=self.loop)) + + host, port = server.sockets[0].getsockname() + + async def do_connect(): + stream = await asyncio.connect(host, port) + stream.write(b'begin\n') + with open(support.TESTFN, 'rb') as fp: + await stream.sendfile(fp) + stream.write(b'end\n') + data = await stream.readline() + self.assertEqual(data, b'done\n') + await stream.aclose() + + self.loop.run_until_complete(do_connect()) + server.close() + self.loop.run_until_complete(server.wait_closed()) + + self.assertEqual([], messages) + + @unittest.skipIf(ssl is None, 'No ssl module') + def test_connect_start_tls(self): + with test_utils.run_test_server(use_ssl=True) as httpd: + # connect without SSL but upgrade to TLS just after + # connection is established + stream = self.loop.run_until_complete( + asyncio.connect(*httpd.address)) + + self.loop.run_until_complete( + stream.start_tls( + sslcontext=test_utils.dummy_ssl_context())) + self._basetest_connect(stream) + if __name__ == '__main__': unittest.main()