Skip to content
Merged
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
30 changes: 27 additions & 3 deletions src/java/net/jpountz/lz4/LZ4BlockInputStream.java
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,8 @@
* {@link InputStream} implementation to decode data written with
* {@link LZ4BlockOutputStream}. This class is not thread-safe and does not
* support {@link #mark(int)}/{@link #reset()}.
* Once a read fails with an {@link IOException}, every later read or skip on
* this stream throws an {@link IOException} as well.
* <p>Use {@link Builder#withAcceptOversizedBlocks(boolean)} only for trusted
* inputs that may contain noncanonical legacy LZ4 blocks. Enabling it restores
* acceptance of blocks whose compressed length is greater than or equal to the
Expand All @@ -57,6 +59,8 @@ public class LZ4BlockInputStream extends FilterInputStream {
private int originalLen;
private int o;
private boolean finished;
// set once refill() fails; the stream is unusable afterwards
private IOException failure;

/**
* Creates a new LZ4 input stream to read from the specified underlying InputStream.
Expand Down Expand Up @@ -187,11 +191,15 @@ public static Builder newBuilder() {

@Override
public int available() throws IOException {
if (failure != null) {
return 0;
}
return originalLen - o;
}

@Override
public int read() throws IOException {
ensureNotFailed();
if (finished) {
return -1;
}
Expand All @@ -207,6 +215,7 @@ public int read() throws IOException {
@Override
public int read(byte[] b, int off, int len) throws IOException {
SafeUtils.checkRange(b, off, len);
ensureNotFailed();
if (finished) {
return -1;
}
Expand All @@ -229,6 +238,7 @@ public int read(byte[] b) throws IOException {

@Override
public long skip(long n) throws IOException {
ensureNotFailed();
if (n <= 0 || finished) {
return 0;
}
Expand All @@ -243,10 +253,24 @@ public long skip(long n) throws IOException {
return skipped;
}

private void ensureNotFailed() throws IOException {
if (failure != null) {
throw new IOException("Stream previously failed", failure);
}
}

private void refill() throws IOException {
// Loop rather than recurse over empty blocks so that a long run of them cannot overflow the stack
while (!readBlock()) {
// empty block with stopOnEmptyBlock == false, continue with the next block
try {
// Loop rather than recurse over empty blocks so that a long run of them cannot overflow the stack
while (!readBlock()) {
// empty block with stopOnEmptyBlock == false, continue with the next block
}
} catch (IOException e) {
failure = e;
throw e;
} catch (RuntimeException e) {
failure = new IOException("Stream is corrupted", e);
throw failure;
}
}

Expand Down
42 changes: 41 additions & 1 deletion src/java/net/jpountz/lz4/LZ4FrameInputStream.java
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,8 @@

/**
* Implementation of the v1.5.1 LZ4 Frame format. This class is NOT thread safe.
* Once a read fails with an {@link IOException}, every later read or skip on
* this stream throws an {@link IOException} as well.
* <p>
* Not Supported:<ul>
* <li>Dependent blocks</li>
Expand Down Expand Up @@ -70,6 +72,8 @@ public class LZ4FrameInputStream extends FilterInputStream {
private boolean anyFrameRead = false;

private LZ4FrameOutputStream.FrameInfo frameInfo = null;
// set once reading a frame header or block fails; the stream is unusable afterwards
private IOException failure = null;

/**
* Creates a new {@link InputStream} that will decompress data using fastest instances of {@link LZ4SafeDecompressor} and {@link XXHash32}.
Expand Down Expand Up @@ -141,6 +145,16 @@ public LZ4FrameInputStream(InputStream in, LZ4SafeDecompressor decompressor, XXH
* @throws IOException On input stream read exception
*/
private boolean nextFrameInfo() throws IOException {
try {
return nextFrameInfo0();
} catch (IOException e) {
throw fail(e);
} catch (RuntimeException e) {
throw fail(new IOException("Stream is corrupted", e));
}
}

private boolean nextFrameInfo0() throws IOException {
while (true) {
int size = 0;
do {
Expand Down Expand Up @@ -286,6 +300,27 @@ private int readInt(InputStream stream) throws IOException {
* @throws IOException
*/
private void readBlock() throws IOException {
try {
readBlock0();
} catch (IOException e) {
throw fail(e);
} catch (RuntimeException e) {
throw fail(new IOException("Stream is corrupted", e));
}
}

private IOException fail(IOException e) {
failure = e;
return e;
}

private void ensureNotFailed() throws IOException {
if (failure != null) {
throw new IOException("Stream previously failed", failure);
}
}

private void readBlock0() throws IOException {
if (frameInfo.isEnabled(LZ4FrameOutputStream.FLG.Bits.CONTENT_CHECKSUM) && streamHash == null) {
// the content checksum hash was released by close()
throw new IOException("Stream closed");
Expand Down Expand Up @@ -366,6 +401,7 @@ private void readBlock() throws IOException {

@Override
public int read() throws IOException {
ensureNotFailed();
while (!firstFrameHeaderRead || buffer.remaining() == 0) {
if (!firstFrameHeaderRead || frameInfo.isFinished()) {
if (firstFrameHeaderRead && readSingleFrame) {
Expand All @@ -385,6 +421,7 @@ public int read(byte[] b, int off, int len) throws IOException {
if ((off < 0) || (len < 0) || notEnoughSpace(b.length - off, len)) {
throw new IndexOutOfBoundsException();
}
ensureNotFailed();
while (!firstFrameHeaderRead || buffer.remaining() == 0) {
if (!firstFrameHeaderRead || frameInfo.isFinished()) {
if (firstFrameHeaderRead && readSingleFrame) {
Expand All @@ -403,6 +440,7 @@ public int read(byte[] b, int off, int len) throws IOException {

@Override
public long skip(long n) throws IOException {
ensureNotFailed();
if (n <= 0) {
return 0;
}
Expand All @@ -424,7 +462,7 @@ public long skip(long n) throws IOException {

@Override
public int available() throws IOException {
if (!firstFrameHeaderRead) {
if (failure != null || !firstFrameHeaderRead) {
return 0;
}
return buffer.remaining();
Expand Down Expand Up @@ -468,6 +506,7 @@ public boolean markSupported() {
* @see #LZ4FrameInputStream(InputStream, LZ4SafeDecompressor, XXHash32, boolean)
*/
public long getExpectedContentSize() throws IOException {
ensureNotFailed();
if (!readSingleFrame) {
throw new UnsupportedOperationException("Operation not permitted when multiple frames can be read");
}
Expand All @@ -486,6 +525,7 @@ public long getExpectedContentSize() throws IOException {
* @throws IOException On input stream read exception
*/
public boolean isExpectedContentSizeDefined() throws IOException {
ensureNotFailed();
if (readSingleFrame) {
if (!firstFrameHeaderRead) {
if (!nextFrameInfo()) {
Expand Down
55 changes: 55 additions & 0 deletions src/test/net/jpountz/lz4/LZ4BlockStreamingTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,8 @@
import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.nio.ByteBuffer;
import java.nio.ByteOrder;
import java.nio.charset.Charset;
import java.util.Arrays;
import java.util.function.Consumer;
Expand Down Expand Up @@ -602,4 +604,57 @@ public void testManyEmptyBlocksSkip() throws IOException {
assertEquals(-1, in.read());
in.close();
}

private static byte[] twoBlockStream() throws IOException {
// first block holds 64 bytes, second block 32 bytes
var baos = new ByteArrayOutputStream();
try (var out = new LZ4BlockOutputStream(baos, 64)) {
for (int i = 0; i < 96; ++i) {
out.write(i);
}
}
return baos.toByteArray();
}

private static void assertFailedStream(LZ4BlockInputStream in) throws IOException {
assertThrows(IOException.class, in::read);
assertThrows(IOException.class, () -> in.read(new byte[16]));
assertThrows(IOException.class, () -> in.skip(1));
assertThrows(IOException.class, in::readAllBytes);
assertThrows(IOException.class, () -> in.readNBytes(16));
assertEquals(0, in.available());
in.close();
}

@Test
public void testFailureIsStickyAfterChecksumMismatch() throws IOException {
byte[] compressed = twoBlockStream();
final int secondBlock = LZ4BlockOutputStream.HEADER_LENGTH
+ net.jpountz.util.SafeUtils.readIntLE(compressed, LZ4BlockOutputStream.MAGIC_LENGTH + 1);
// corrupt the checksum of the second block
compressed[secondBlock + LZ4BlockOutputStream.MAGIC_LENGTH + 9] ^= 1;

LZ4BlockInputStream in = lz4BlockInputStreamBuilder().build(new ByteArrayInputStream(compressed));
byte[] first = new byte[64];
assertEquals(64, in.readNBytes(first, 0, 64));
var e = assertThrows(IOException.class, in::read);
assertEquals("Stream is corrupted", e.getMessage());
assertFailedStream(in);
}

@Test
public void testFailureIsStickyAfterInvalidOriginalLength() throws IOException {
for (int originalLen : new int[] {-1, Integer.MAX_VALUE}) {
byte[] compressed = twoBlockStream();
final int secondBlock = LZ4BlockOutputStream.HEADER_LENGTH
+ net.jpountz.util.SafeUtils.readIntLE(compressed, LZ4BlockOutputStream.MAGIC_LENGTH + 1);
ByteBuffer.wrap(compressed).order(ByteOrder.LITTLE_ENDIAN).putInt(secondBlock + LZ4BlockOutputStream.MAGIC_LENGTH + 5, originalLen);

LZ4BlockInputStream in = lz4BlockInputStreamBuilder().build(new ByteArrayInputStream(compressed));
byte[] first = new byte[64];
assertEquals(64, in.readNBytes(first, 0, 64));
assertThrows(IOException.class, () -> in.read(new byte[16]));
assertFailedStream(in);
}
}
}
118 changes: 116 additions & 2 deletions src/test/net/jpountz/lz4/LZ4FrameIOStreamTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -48,8 +48,6 @@
import java.util.List;
import java.util.Random;

import net.jpountz.xxhash.XXHashFactory;

/**
*
*/
Expand Down Expand Up @@ -868,6 +866,122 @@ public void testAvailable() throws IOException {
}
}

private static byte[] compressFrame(byte[] data, LZ4FrameOutputStream.FLG.Bits... bits) throws IOException {
final ByteArrayOutputStream baos = new ByteArrayOutputStream();
try (OutputStream os = new LZ4FrameOutputStream(baos, LZ4FrameOutputStream.BLOCKSIZE.SIZE_64KB, bits)) {
os.write(data);
}
return baos.toByteArray();
}

private static byte[] concat(byte[]... arrays) {
final ByteArrayOutputStream baos = new ByteArrayOutputStream();
for (byte[] a : arrays) {
baos.write(a, 0, a.length);
}
return baos.toByteArray();
}

private static void assertFailedStream(LZ4FrameInputStream is) throws IOException {
Assert.assertThrows(IOException.class, is::read);
Assert.assertThrows(IOException.class, () -> is.read(new byte[16]));
Assert.assertThrows(IOException.class, () -> is.skip(1));
Assert.assertThrows(IOException.class, is::readAllBytes);
Assert.assertThrows(IOException.class, () -> is.readNBytes(16));
Assert.assertEquals(0, is.available());
is.close();
}

@Test
public void testFailureIsStickyAfterDescriptorHashMismatch() throws IOException {
final byte[] frame1 = compressFrame("hello".getBytes("UTF-8"), LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE);
final byte[] frame2 = compressFrame("world".getBytes("UTF-8"), LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE);
// magic (4), FLG, BD, then the header checksum
frame2[6] ^= 1;

final LZ4FrameInputStream is = new LZ4FrameInputStream(new ByteArrayInputStream(concat(frame1, frame2)));
final byte[] first = new byte[5];
Assert.assertEquals(5, is.readNBytes(first, 0, 5));
Assert.assertArrayEquals("hello".getBytes("UTF-8"), first);
final IOException e = Assert.assertThrows(IOException.class, is::read);
Assert.assertEquals(LZ4FrameInputStream.DESCRIPTOR_HASH_MISMATCH, e.getMessage());
// must not go on to return the data of the rejected frame
assertFailedStream(is);
}

@Test
public void testFailureIsStickyAfterSkippableFrameAndCorruptHeader() throws IOException {
final byte[] skippable = new byte[8];
ByteBuffer.wrap(skippable).order(ByteOrder.LITTLE_ENDIAN).putInt(LZ4FrameInputStream.MAGIC_SKIPPABLE_BASE).putInt(0);
final byte[] frame = compressFrame("hello".getBytes("UTF-8"), LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE);
frame[6] ^= 1;

final LZ4FrameInputStream is = new LZ4FrameInputStream(new ByteArrayInputStream(concat(skippable, frame)));
Assert.assertThrows(IOException.class, is::read);
assertFailedStream(is);
}

@Test
public void testFailureIsStickyAfterBlockChecksumMismatch() throws IOException {
final byte[] frame = compressFrame("hello".getBytes("UTF-8"), LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE,
LZ4FrameOutputStream.FLG.Bits.BLOCK_CHECKSUM);
// the frame ends with the block checksum (4) and the end mark (4)
frame[frame.length - 5] ^= 1;

final LZ4FrameInputStream is = new LZ4FrameInputStream(new ByteArrayInputStream(frame));
final IOException e = Assert.assertThrows(IOException.class, is::read);
Assert.assertEquals(LZ4FrameInputStream.BLOCK_HASH_MISMATCH, e.getMessage());
// must not silently drop the block and continue
assertFailedStream(is);
}

@Test
public void testMalformedFirstHeaderThrowsIOException() throws IOException {
final byte[] frame = compressFrame("hello".getBytes("UTF-8"), LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE);
// set the reserved bit 0 in FLG
frame[4] |= 1;

// the header is read lazily, so construction must not fail
final LZ4FrameInputStream is = new LZ4FrameInputStream(new ByteArrayInputStream(frame));
Assert.assertThrows(IOException.class, is::read);
assertFailedStream(is);
}

@Test
public void testLinkedBlocksFirstHeaderThrowsIOException() throws IOException {
final byte[] frame = compressFrame("hello".getBytes("UTF-8"), LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE);
// version 01, BLOCK_INDEPENDENCE cleared (linked blocks), with a valid descriptor checksum
frame[4] = 0x40;
frame[6] = (byte) ((XXHashFactory.fastestInstance().hash32().hash(frame, 4, 2, 0) >> 8) & 0xFF);

// the header is read lazily, so construction must not fail
final LZ4FrameInputStream is = new LZ4FrameInputStream(new ByteArrayInputStream(frame), true);
assertInvalidDescriptor(Assert.assertThrows(IOException.class, is::isExpectedContentSizeDefined));
Assert.assertThrows(IOException.class, is::getExpectedContentSize);
assertFailedStream(is);
}

@Test
public void testFailureIsStickyForExpectedContentSize() throws IOException {
final byte[] skippable = new byte[8];
ByteBuffer.wrap(skippable).order(ByteOrder.LITTLE_ENDIAN).putInt(LZ4FrameInputStream.MAGIC_SKIPPABLE_BASE).putInt(0);
final ByteArrayOutputStream baos = new ByteArrayOutputStream();
try (OutputStream os = new LZ4FrameOutputStream(baos, LZ4FrameOutputStream.BLOCKSIZE.SIZE_64KB, 5L,
LZ4FrameOutputStream.FLG.Bits.BLOCK_INDEPENDENCE, LZ4FrameOutputStream.FLG.Bits.CONTENT_SIZE)) {
os.write("hello".getBytes("UTF-8"));
}
final byte[] frame = baos.toByteArray();
// magic (4), FLG, BD, content size (8), then the header checksum
frame[14] ^= 1;

final LZ4FrameInputStream is = new LZ4FrameInputStream(new ByteArrayInputStream(concat(skippable, frame)), true);
final IOException e = Assert.assertThrows(IOException.class, is::getExpectedContentSize);
Assert.assertEquals(LZ4FrameInputStream.DESCRIPTOR_HASH_MISMATCH, e.getMessage());
// must not report the content size of the rejected frame
Assert.assertThrows(IOException.class, is::isExpectedContentSizeDefined);
assertFailedStream(is);
}

private static byte[] frameHeader(int flg, int bd) {
final byte[] descriptor = new byte[]{(byte) flg, (byte) bd};
final byte hc = (byte) ((XXHashFactory.fastestInstance().hash32().hash(descriptor, 0, descriptor.length, 0) >> 8) & 0xFF);
Expand Down
Loading