]> git.ipfire.org Git - thirdparty/Python/cpython.git/commitdiff
gh-150621: avoid quadratic bytes slicing in `asyncio.protocols._feed_data_to_buffered...
authorTimofei <128279579+deadlovelll@users.noreply.github.com>
Sat, 25 Jul 2026 13:14:49 +0000 (16:14 +0300)
committerGitHub <noreply@github.com>
Sat, 25 Jul 2026 13:14:49 +0000 (13:14 +0000)
Lib/asyncio/protocols.py
Lib/test/test_asyncio/test_protocols.py

index 09987b164c66e57f8bd12190899a86947cd065b6..c67a26c9e2986cc4b004dd891f4e6fb535dc970c 100644 (file)
@@ -199,6 +199,7 @@ class SubprocessProtocol(BaseProtocol):
 
 def _feed_data_to_buffered_proto(proto, data):
     data_len = len(data)
+    start = 0
     while data_len:
         buf = proto.get_buffer(data_len)
         buf_len = len(buf)
@@ -206,11 +207,11 @@ def _feed_data_to_buffered_proto(proto, data):
             raise RuntimeError('get_buffer() returned an empty buffer')
 
         if buf_len >= data_len:
-            buf[:data_len] = data
+            buf[:data_len] = data[start:] if start else data
             proto.buffer_updated(data_len)
             return
         else:
-            buf[:buf_len] = data[:buf_len]
+            buf[:buf_len] = data[start:start + buf_len]
             proto.buffer_updated(buf_len)
-            data = data[buf_len:]
-            data_len = len(data)
+            start += buf_len
+            data_len -= buf_len
index 643199962b8af066dd1ee68170012a05d1d36a31..38f1e3fba90cfdea5d92c3a1bf254e4ed6d457c8 100644 (file)
@@ -2,6 +2,7 @@ import unittest
 from unittest import mock
 
 import asyncio
+from asyncio import protocols
 
 
 def tearDownModule():
@@ -63,5 +64,32 @@ class ProtocolsAbsTests(unittest.TestCase):
         self.assertNotHasAttr(sp, '__dict__')
 
 
-if __name__ == '__main__':
+class FeedDataToBufferedProtoTests(unittest.TestCase):
+    def _make_proto(self, bufsize):
+        received = bytearray()
+        buf = bytearray(bufsize)
+
+        class P(asyncio.BufferedProtocol):
+            def get_buffer(self, sizehint):
+                return buf
+
+            def buffer_updated(self, nbytes):
+                received.extend(buf[:nbytes])
+
+        return P(), received
+
+    def test_large_multi_iteration(self):
+        proto, received = self._make_proto(64)
+        data = bytes(range(256)) * 16
+        protocols._feed_data_to_buffered_proto(proto, data)
+        self.assertEqual(bytes(received), data)
+
+    def test_memoryview_input(self):
+        proto, received = self._make_proto(64)
+        payload = b"y" * 200
+        protocols._feed_data_to_buffered_proto(proto, memoryview(payload))
+        self.assertEqual(bytes(received), payload)
+
+
+if __name__ == "__main__":
     unittest.main()