Skip to content

Commit a7e4921

Browse files
miss-islingtonkn1g78kumaraditya303
authored
[3.14] gh-152431: update StreamReader transport after StreamWriter.start_tls() (GH-152432) (#154629)
gh-152431: update StreamReader transport after StreamWriter.start_tls() (GH-152432) (cherry picked from commit 7671ee1) Co-authored-by: Xuyang Zhang <119476662+kn1g78@users.noreply.github.com> Co-authored-by: Kumar Aditya <kumaraditya@python.org>
1 parent e8a93c7 commit a7e4921

3 files changed

Lines changed: 37 additions & 0 deletions

File tree

Lib/asyncio/streams.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -217,6 +217,9 @@ def _replace_transport(self, transport):
217217
loop = self._loop
218218
self._transport = transport
219219
self._over_ssl = transport.get_extra_info('sslcontext') is not None
220+
reader = self._stream_reader
221+
if reader is not None:
222+
reader._replace_transport(transport)
220223

221224
def connection_made(self, transport):
222225
if self._reject_connection:
@@ -477,6 +480,10 @@ def set_transport(self, transport):
477480
assert self._transport is None, 'Transport already set'
478481
self._transport = transport
479482

483+
def _replace_transport(self, transport):
484+
assert self._transport is not None, 'Transport not set'
485+
self._transport = transport
486+
480487
def _maybe_resume_transport(self):
481488
if self._paused and len(self._buffer) <= self._limit:
482489
self._paused = False

Lib/test/test_asyncio/test_streams.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -861,6 +861,34 @@ async def run_test():
861861
self.loop.run_until_complete(run_test())
862862
self.assertEqual(messages, [])
863863

864+
def test_streamwriter_start_tls_updates_reader_transport(self):
865+
reader = asyncio.StreamReader(loop=self.loop)
866+
protocol = asyncio.StreamReaderProtocol(reader, loop=self.loop)
867+
old_transport = mock.Mock()
868+
old_transport.get_extra_info.return_value = None
869+
old_transport.is_closing.return_value = False
870+
protocol.connection_made(old_transport)
871+
872+
writer = asyncio.StreamWriter(old_transport, protocol, reader, self.loop)
873+
874+
ssl_context = mock.sentinel.ssl_context
875+
new_transport = mock.Mock()
876+
new_transport.get_extra_info.return_value = ssl_context
877+
self.loop.start_tls = mock.AsyncMock(return_value=new_transport)
878+
879+
self.loop.run_until_complete(writer.start_tls(ssl_context))
880+
881+
self.loop.start_tls.assert_awaited_once_with(
882+
old_transport, protocol, ssl_context,
883+
server_side=False, server_hostname=None,
884+
ssl_handshake_timeout=None,
885+
ssl_shutdown_timeout=None,
886+
)
887+
self.assertIs(writer.transport, new_transport)
888+
self.assertIs(protocol._transport, new_transport)
889+
self.assertIs(reader._transport, new_transport)
890+
self.assertTrue(protocol._over_ssl)
891+
864892
def test_streamreader_constructor_without_loop(self):
865893
with self.assertRaisesRegex(RuntimeError, 'no current event loop'):
866894
asyncio.StreamReader()
Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
Fix ``asyncio.StreamWriter.start_tls()`` to keep the linked
2+
``StreamReader`` transport in sync with the upgraded transport.

0 commit comments

Comments
 (0)