]> git.ipfire.org Git - thirdparty/Python/cpython.git/commitdiff
gh-152431: update StreamReader transport after StreamWriter.start_tls() (#152432)
authorXuyang Zhang <119476662+kn1g78@users.noreply.github.com>
Sat, 11 Jul 2026 13:03:23 +0000 (21:03 +0800)
committerGitHub <noreply@github.com>
Sat, 11 Jul 2026 13:03:23 +0000 (13:03 +0000)
Co-authored-by: Kumar Aditya <kumaraditya@python.org>
Lib/asyncio/streams.py
Lib/test/test_asyncio/test_streams.py
Misc/NEWS.d/next/Library/2026-06-28-00-00-00.gh-issue-152431.Ja3K9m.rst [new file with mode: 0644]

index d6076465f79366fa85f4654b8dde64fa118395d9..f5c4f0b0c3297ba9add38f9e3c8e519d78daf295 100644 (file)
@@ -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
index fb774895a7e52a545771f5233839dd53aada69cc..911087a128f9713d293c5e58c790e6a3d60689f9 100644 (file)
@@ -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 (file)
index 0000000..bda2dfd
--- /dev/null
@@ -0,0 +1,2 @@
+Fix ``asyncio.StreamWriter.start_tls()`` to keep the linked
+``StreamReader`` transport in sync with the upgraded transport.