From c919487fc34470e4e674e489da58307e181f4e8b Mon Sep 17 00:00:00 2001 From: CAOShurong <170531907+CAOShurong@users.noreply.github.com> Date: Fri, 9 Oct 2026 14:06:01 +0800 Subject: [PATCH] Fix stream reader readinto position accounting --- c-ext/compressionreader.c | 4 +- tests/test_compressor_stream_reader.py | 95 ++++++++++++++++++++++++++ 2 files changed, 97 insertions(+), 2 deletions(-) diff --git a/c-ext/compressionreader.c b/c-ext/compressionreader.c index b14a63b1..d9027680 100644 --- a/c-ext/compressionreader.c +++ b/c-ext/compressionreader.c @@ -593,7 +593,7 @@ static PyObject *compressionreader_readinto(ZstdCompressionReader *self, zresult = ZSTD_compressStream2(self->compressor->cctx, &output, &self->input, ZSTD_e_end); - self->bytesCompressed += self->output.pos - oldPos; + self->bytesCompressed += output.pos - oldPos; if (ZSTD_isError(zresult)) { PyErr_Format(ZstdError, "error ending compression stream: %s", @@ -703,7 +703,7 @@ static PyObject *compressionreader_readinto1(ZstdCompressionReader *self, zresult = ZSTD_compressStream2(self->compressor->cctx, &output, &self->input, ZSTD_e_end); - self->bytesCompressed += self->output.pos - oldPos; + self->bytesCompressed += output.pos - oldPos; if (ZSTD_isError(zresult)) { PyErr_Format(ZstdError, "error ending compression stream: %s", diff --git a/tests/test_compressor_stream_reader.py b/tests/test_compressor_stream_reader.py index 6c1fca21..b247d1e2 100644 --- a/tests/test_compressor_stream_reader.py +++ b/tests/test_compressor_stream_reader.py @@ -246,6 +246,101 @@ def test_readinto1(self): self.assertEqual(b[0 : len(foo)], foo) self.assertEqual(source._read_count, 4) + def test_readinto_tell(self): + data = bytes(range(256)) * 2048 + cases = ( + (b"foo", False), + (b"foo", True), + (b"", True), + (data, False), + (data, True), + ) + + for method in ("readinto", "readinto1"): + for data, stream in cases: + for size in (1, 17, 1024): + with self.subTest( + method=method, + stream=stream, + source_size=len(data), + buffer_size=size, + ): + cctx = zstd.ZstdCompressor(write_checksum=True) + source = io.BytesIO(data) if stream else data + dest = bytearray(size) + chunks = [] + total = 0 + + with cctx.stream_reader( + source, size=len(data) + ) as reader: + self.assertEqual(reader.tell(), 0) + while True: + count = getattr(reader, method)(dest) + total += count + chunks.append(bytes(dest[:count])) + self.assertEqual(reader.tell(), total) + if not count: + break + + self.assertEqual(getattr(reader, method)(dest), 0) + self.assertEqual(reader.tell(), total) + + frame = b"".join(chunks) + self.assertEqual( + zstd.ZstdDecompressor().decompress( + frame, max_output_size=max(len(data), 1) + ), + data, + ) + + def test_readinto_tell_mixed_reads(self): + for method in ("readinto", "readinto1"): + with self.subTest(method=method): + cctx = zstd.ZstdCompressor() + data = b"foo" * 1024 + dest = bytearray(b"!" * 1026) + view = memoryview(dest)[1:-1] + chunks = [] + + with cctx.stream_reader( + io.BytesIO(data), size=len(data) + ) as reader: + chunks.append(reader.read(1)) + chunks.append(reader.read1(2)) + total = sum(map(len, chunks)) + self.assertEqual(reader.tell(), total) + + while True: + count = getattr(reader, method)(view) + total += count + chunks.append(bytes(view[:count])) + self.assertEqual(reader.tell(), total) + self.assertEqual(dest[0], ord("!")) + self.assertEqual(dest[-1], ord("!")) + if not count: + break + + self.assertEqual(reader.read(1), b"") + self.assertEqual(reader.read1(1), b"") + self.assertEqual(reader.tell(), total) + + self.assertEqual( + zstd.ZstdDecompressor().decompress(b"".join(chunks)), data + ) + + def test_readinto_tell_invalid_buffer(self): + for method in ("readinto", "readinto1"): + with self.subTest(method=method): + reader = zstd.ZstdCompressor().stream_reader(b"foo") + with self.assertRaises((TypeError, BufferError, ValueError)): + getattr(reader, method)(b"read only") + self.assertEqual(reader.tell(), 0) + reader.close() + with self.assertRaisesRegex(ValueError, "stream is closed"): + getattr(reader, method)(bytearray(1)) + self.assertEqual(reader.tell(), 0) + def test_read1(self): cctx = zstd.ZstdCompressor() foo = b"".join(cctx.read_to_iter(io.BytesIO(b"foo")))