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)
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
from unittest import mock
import asyncio
+from asyncio import protocols
def tearDownModule():
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()