Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions Doc/library/asyncio-stream.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
========
Expand Down
74 changes: 71 additions & 3 deletions Lib/asyncio/streams.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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.
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
81 changes: 81 additions & 0 deletions Lib/test/test_asyncio/test_streams.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()