From: Xuyang Zhang <119476662+kn1g78@users.noreply.github.com> Date: Sat, 11 Jul 2026 13:03:23 +0000 (+0800) Subject: gh-152431: update StreamReader transport after StreamWriter.start_tls() (#152432) X-Git-Url: http://git.ipfire.org/gitweb/index.cgi?a=commitdiff_plain;h=7671ee1eba3c7c13747761b35b4b9d4166a4670a;p=thirdparty%2FPython%2Fcpython.git gh-152431: update StreamReader transport after StreamWriter.start_tls() (#152432) Co-authored-by: Kumar Aditya --- diff --git a/Lib/asyncio/streams.py b/Lib/asyncio/streams.py index d6076465f793..f5c4f0b0c329 100644 --- a/Lib/asyncio/streams.py +++ b/Lib/asyncio/streams.py @@ -216,6 +216,9 @@ class StreamReaderProtocol(FlowControlMixin, protocols.Protocol): def _replace_transport(self, transport): self._transport = transport self._over_ssl = transport.get_extra_info('sslcontext') is not None + reader = self._stream_reader + if reader is not None: + reader._replace_transport(transport) def connection_made(self, transport): if self._reject_connection: @@ -473,6 +476,10 @@ class StreamReader: assert self._transport is None, 'Transport already set' self._transport = transport + def _replace_transport(self, transport): + assert self._transport is not None, 'Transport not set' + self._transport = transport + def _maybe_resume_transport(self): if self._paused and len(self._buffer) <= self._limit: self._paused = False diff --git a/Lib/test/test_asyncio/test_streams.py b/Lib/test/test_asyncio/test_streams.py index fb774895a7e5..911087a128f9 100644 --- a/Lib/test/test_asyncio/test_streams.py +++ b/Lib/test/test_asyncio/test_streams.py @@ -861,6 +861,34 @@ class StreamTests(test_utils.TestCase): self.loop.run_until_complete(run_test()) self.assertEqual(messages, []) + def test_streamwriter_start_tls_updates_reader_transport(self): + reader = asyncio.StreamReader(loop=self.loop) + protocol = asyncio.StreamReaderProtocol(reader, loop=self.loop) + old_transport = mock.Mock() + old_transport.get_extra_info.return_value = None + old_transport.is_closing.return_value = False + protocol.connection_made(old_transport) + + writer = asyncio.StreamWriter(old_transport, protocol, reader, self.loop) + + ssl_context = mock.sentinel.ssl_context + new_transport = mock.Mock() + new_transport.get_extra_info.return_value = ssl_context + self.loop.start_tls = mock.AsyncMock(return_value=new_transport) + + self.loop.run_until_complete(writer.start_tls(ssl_context)) + + self.loop.start_tls.assert_awaited_once_with( + old_transport, protocol, ssl_context, + server_side=False, server_hostname=None, + ssl_handshake_timeout=None, + ssl_shutdown_timeout=None, + ) + self.assertIs(writer.transport, new_transport) + self.assertIs(protocol._transport, new_transport) + self.assertIs(reader._transport, new_transport) + self.assertTrue(protocol._over_ssl) + def test_streamreader_constructor_without_loop(self): with self.assertRaisesRegex(RuntimeError, 'no current event loop'): asyncio.StreamReader() diff --git a/Misc/NEWS.d/next/Library/2026-06-28-00-00-00.gh-issue-152431.Ja3K9m.rst b/Misc/NEWS.d/next/Library/2026-06-28-00-00-00.gh-issue-152431.Ja3K9m.rst new file mode 100644 index 000000000000..bda2dfd7cd9d --- /dev/null +++ b/Misc/NEWS.d/next/Library/2026-06-28-00-00-00.gh-issue-152431.Ja3K9m.rst @@ -0,0 +1,2 @@ +Fix ``asyncio.StreamWriter.start_tls()`` to keep the linked +``StreamReader`` transport in sync with the upgraded transport.