Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions c-ext/compressionreader.c
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
95 changes: 95 additions & 0 deletions tests/test_compressor_stream_reader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")))
Expand Down