diff --git a/src/java/net/jpountz/lz4/LZ4BlockInputStream.java b/src/java/net/jpountz/lz4/LZ4BlockInputStream.java
index 953f5a4..3408d64 100644
--- a/src/java/net/jpountz/lz4/LZ4BlockInputStream.java
+++ b/src/java/net/jpountz/lz4/LZ4BlockInputStream.java
@@ -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.
*
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
@@ -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.
@@ -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;
}
@@ -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;
}
@@ -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;
}
@@ -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;
}
}
diff --git a/src/java/net/jpountz/lz4/LZ4FrameInputStream.java b/src/java/net/jpountz/lz4/LZ4FrameInputStream.java
index d1aaa74..3bd0fec 100644
--- a/src/java/net/jpountz/lz4/LZ4FrameInputStream.java
+++ b/src/java/net/jpountz/lz4/LZ4FrameInputStream.java
@@ -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.
*
* Not Supported:
* - Dependent blocks
@@ -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}.
@@ -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 {
@@ -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");
@@ -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) {
@@ -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) {
@@ -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;
}
@@ -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();
@@ -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");
}
@@ -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()) {
diff --git a/src/test/net/jpountz/lz4/LZ4BlockStreamingTest.java b/src/test/net/jpountz/lz4/LZ4BlockStreamingTest.java
index 46c94be..d9a2576 100644
--- a/src/test/net/jpountz/lz4/LZ4BlockStreamingTest.java
+++ b/src/test/net/jpountz/lz4/LZ4BlockStreamingTest.java
@@ -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;
@@ -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);
+ }
+ }
}
diff --git a/src/test/net/jpountz/lz4/LZ4FrameIOStreamTest.java b/src/test/net/jpountz/lz4/LZ4FrameIOStreamTest.java
index 3de8f2b..2cc18ff 100644
--- a/src/test/net/jpountz/lz4/LZ4FrameIOStreamTest.java
+++ b/src/test/net/jpountz/lz4/LZ4FrameIOStreamTest.java
@@ -48,8 +48,6 @@
import java.util.List;
import java.util.Random;
-import net.jpountz.xxhash.XXHashFactory;
-
/**
*
*/
@@ -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);