diff --git a/src/xz/file.py b/src/xz/file.py index 7d96993..271ca3b 100644 --- a/src/xz/file.py +++ b/src/xz/file.py @@ -178,7 +178,12 @@ def _init_parse(self) -> None: self._append(streams.pop()) def _create_fileobj(self) -> XZStream: - stream_pos = sum(len(stream.fileobj) for stream in self._fileobjs.values()) + last_stream = self._last_stream + stream_pos = ( + last_stream.fileobj.start + len(last_stream.fileobj) + if last_stream is not None + else 0 + ) return XZStream( IOProxy( self.fileobj, diff --git a/tests/unit/test_file.py b/tests/unit/test_file.py index 09379b8..cf0d25b 100644 --- a/tests/unit/test_file.py +++ b/tests/unit/test_file.py @@ -1,5 +1,6 @@ from collections.abc import Callable from io import SEEK_END, SEEK_SET, BytesIO, UnsupportedOperation +from lzma import compress import os from pathlib import Path from typing import cast @@ -494,6 +495,71 @@ def test_write_with_mode( assert value == expected_value +@pytest.mark.parametrize("from_file", [False, True]) +@pytest.mark.parametrize( + ["payloads", "padding_size", "trailing_padding_size"], + [ + ((b"first", b"second"), 0, 0), + ((b"first", b"second"), 4, 0), + ((b"first", b"second"), 12, 0), + ((b"first", b"second"), 64, 8), + ((b"", b"first", b"second"), 0, 0), + ((b"first", b"", b"second"), 0, 0), + ((b"first", b"second", b""), 4, 8), + ], +) +def test_append_stream_after_existing_gaps( + payloads: tuple[bytes, ...], + padding_size: int, + trailing_padding_size: int, + from_file: bool, + tmp_path: Path, +) -> None: + original = bytes(padding_size).join(compress(data) for data in payloads) + filename: Path | BytesIO + if from_file: + filename = tmp_path / "archive.xz" + filename.write_bytes(original + bytes(trailing_padding_size)) + else: + filename = BytesIO(original + bytes(trailing_padding_size)) + filename.seek(7) + + with XZFile(filename, "r+") as xzfile: + assert xzfile.read() == b"".join(payloads) + xzfile.seek(0, SEEK_END) + xzfile.change_stream() + xzfile.change_stream() # an unwritten stream is replaced + xzfile.write(b"third") + + if from_file: + edited = cast("Path", filename).read_bytes() + else: + edited = cast("BytesIO", filename).getvalue() + + if payloads[-1]: + assert edited.startswith(original) + with XZFile(BytesIO(edited)) as xzfile: + assert xzfile.read() == b"".join(payloads) + b"third" + + +def test_append_stream_after_truncate() -> None: + payloads = (b"first", b"second", b"third", b"fourth") + encoded = [compress(data) for data in payloads] + filename = BytesIO(bytes(4).join(encoded) + bytes(8)) + expected_prefix = bytes(4).join(encoded[:-1]) + expected_data = b"".join(payloads[:-1]) + + with XZFile(filename, "r+") as xzfile: + xzfile.truncate(len(expected_data)) + xzfile.seek(0, SEEK_END) + xzfile.change_stream() + xzfile.write(b"replacement") + + assert filename.getvalue().startswith(expected_prefix) + with XZFile(BytesIO(filename.getvalue())) as xzfile: + assert xzfile.read() == expected_data + b"replacement" + + # # check / filters / preset changes #