diff --git a/CHANGES.txt b/CHANGES.txt index 4182cc19d2..608d8f8546 100644 --- a/CHANGES.txt +++ b/CHANGES.txt @@ -12,6 +12,7 @@ Merged from 2.2: * (Hadoop) fix splits calculation (CASSANDRA-10640) * (Hadoop) ensure that Cluster instances are always closed (CASSANDRA-10058) Merged from 2.1: + * Fix CompressedInputStream for proper cleanup (CASSANDRA-10012) * (cqlsh) Support counters in COPY commands (CASSANDRA-9043) * Try next replica if not possible to connect to primary replica on ColumnFamilyRecordReader (CASSANDRA-2388) diff --git a/src/java/org/apache/cassandra/streaming/compress/CompressedInputStream.java b/src/java/org/apache/cassandra/streaming/compress/CompressedInputStream.java index ccd0ac5094..56dc63abc7 100644 --- a/src/java/org/apache/cassandra/streaming/compress/CompressedInputStream.java +++ b/src/java/org/apache/cassandra/streaming/compress/CompressedInputStream.java @@ -45,7 +45,7 @@ public class CompressedInputStream extends InputStream private final Supplier crcCheckChanceSupplier; // uncompressed bytes - private byte[] buffer; + private final byte[] buffer; // offset from the beginning of the buffer protected long bufferOffset = 0; @@ -63,6 +63,8 @@ public class CompressedInputStream extends InputStream private long totalCompressedBytesRead; + private Thread readerThread; + /** * @param source Input source to read compressed data from * @param info Compression info @@ -73,10 +75,11 @@ public class CompressedInputStream extends InputStream this.checksum = checksumType.newInstance(); this.buffer = new byte[info.parameters.chunkLength()]; // buffer is limited to store up to 1024 chunks - this.dataBuffer = new ArrayBlockingQueue(Math.min(info.chunks.length, 1024)); + this.dataBuffer = new ArrayBlockingQueue<>(Math.min(info.chunks.length, 1024)); this.crcCheckChanceSupplier = crcCheckChanceSupplier; - new Thread(new Reader(source, info, dataBuffer)).start(); + readerThread = new Thread(new Reader(source, info, dataBuffer)); + readerThread.start(); } public int read() throws IOException @@ -135,7 +138,7 @@ public class CompressedInputStream extends InputStream return totalCompressedBytesRead; } - static class Reader extends WrappedRunnable + class Reader extends WrappedRunnable { private final InputStream source; private final Iterator chunks; @@ -151,7 +154,7 @@ public class CompressedInputStream extends InputStream protected void runMayThrow() throws Exception { byte[] compressedWithCRC; - while (chunks.hasNext()) + while (!Thread.currentThread().isInterrupted() && chunks.hasNext()) { CompressionMetadata.Chunk chunk = chunks.next(); @@ -161,16 +164,43 @@ public class CompressedInputStream extends InputStream int bufferRead = 0; while (bufferRead < readLength) { - int r = source.read(compressedWithCRC, bufferRead, readLength - bufferRead); - if (r < 0) + int r; + try + { + r = source.read(compressedWithCRC, bufferRead, readLength - bufferRead); + if (r < 0) + { + dataBuffer.put(POISON_PILL); + return; // throw exception where we consume dataBuffer + } + } + catch (IOException e) { dataBuffer.put(POISON_PILL); - return; // throw exception where we consume dataBuffer + throw e; } bufferRead += r; } dataBuffer.put(compressedWithCRC); } + synchronized(CompressedInputStream.this) + { + readerThread = null; + } } } + + @Override + public void close() throws IOException + { + synchronized(this) + { + if (readerThread != null) + { + readerThread.interrupt(); + readerThread = null; + } + } + } + } diff --git a/src/java/org/apache/cassandra/streaming/compress/CompressedStreamReader.java b/src/java/org/apache/cassandra/streaming/compress/CompressedStreamReader.java index 8f53832cb7..29882bc155 100644 --- a/src/java/org/apache/cassandra/streaming/compress/CompressedStreamReader.java +++ b/src/java/org/apache/cassandra/streaming/compress/CompressedStreamReader.java @@ -114,6 +114,10 @@ public class CompressedStreamReader extends StreamReader else throw Throwables.propagate(e); } + finally + { + cis.close(); + } } @Override diff --git a/test/unit/org/apache/cassandra/streaming/compression/CompressedInputStreamTest.java b/test/unit/org/apache/cassandra/streaming/compression/CompressedInputStreamTest.java index 5646592bd0..2162e32e27 100644 --- a/test/unit/org/apache/cassandra/streaming/compression/CompressedInputStreamTest.java +++ b/test/unit/org/apache/cassandra/streaming/compression/CompressedInputStreamTest.java @@ -19,6 +19,8 @@ package org.apache.cassandra.streaming.compression; import java.io.*; import java.util.*; +import java.util.concurrent.SynchronousQueue; +import java.util.concurrent.TimeUnit; import org.junit.Test; import org.apache.cassandra.db.ClusteringComparator; @@ -55,6 +57,53 @@ public class CompressedInputStreamTest { testCompressedReadWith(new long[]{1L, 122L, 123L, 124L, 456L}, true); } + + /** + * Test CompressedInputStream not hang when closed while reading + * @throws IOException + */ + @Test(expected = EOFException.class) + public void testClose() throws IOException + { + CompressionParams param = CompressionParams.snappy(32); + CompressionMetadata.Chunk[] chunks = {new CompressionMetadata.Chunk(0, 100)}; + final SynchronousQueue blocker = new SynchronousQueue<>(); + InputStream blockingInput = new InputStream() + { + @Override + public int read() throws IOException + { + try + { + // 10 second cut off not to stop other test in case + return Objects.requireNonNull(blocker.poll(10, TimeUnit.SECONDS)); + } + catch (InterruptedException e) + { + throw new IOException("Interrupted as expected", e); + } + } + }; + CompressionInfo info = new CompressionInfo(chunks, param); + try (CompressedInputStream cis = new CompressedInputStream(blockingInput, info, ChecksumType.CRC32, () -> 1.0)) + { + new Thread(new Runnable() + { + @Override + public void run() + { + try + { + cis.close(); + } + catch (Exception ignore) {} + } + }).start(); + // block here + cis.read(); + } + } + /** * @param valuesToCheck array of longs of range(0-999) * @throws Exception @@ -68,17 +117,19 @@ public class CompressedInputStreamTest Descriptor desc = Descriptor.fromFilename(tmp.getAbsolutePath()); MetadataCollector collector = new MetadataCollector(new ClusteringComparator(BytesType.instance)); CompressionParams param = CompressionParams.snappy(32); - CompressedSequentialWriter writer = new CompressedSequentialWriter(tmp, desc.filenameFor(Component.COMPRESSION_INFO), param, collector); Map index = new HashMap(); - for (long l = 0L; l < 1000; l++) + try (CompressedSequentialWriter writer = new CompressedSequentialWriter(tmp, desc.filenameFor(Component.COMPRESSION_INFO), param, collector)) { - index.put(l, writer.position()); - writer.writeLong(l); + for (long l = 0L; l < 1000; l++) + { + index.put(l, writer.position()); + writer.writeLong(l); + } + writer.finish(); } - writer.finish(); CompressionMetadata comp = CompressionMetadata.create(tmp.getAbsolutePath()); - List> sections = new ArrayList>(); + List> sections = new ArrayList<>(); for (long l : valuesToCheck) { long position = index.get(l); @@ -97,14 +148,15 @@ public class CompressedInputStreamTest size += (c.length + 4); // 4bytes CRC byte[] toRead = new byte[size]; - RandomAccessFile f = new RandomAccessFile(tmp, "r"); - int pos = 0; - for (CompressionMetadata.Chunk c : chunks) + try (RandomAccessFile f = new RandomAccessFile(tmp, "r")) { - f.seek(c.offset); - pos += f.read(toRead, pos, c.length + 4); + int pos = 0; + for (CompressionMetadata.Chunk c : chunks) + { + f.seek(c.offset); + pos += f.read(toRead, pos, c.length + 4); + } } - f.close(); if (testTruncate) { @@ -117,13 +169,15 @@ public class CompressedInputStreamTest CompressionInfo info = new CompressionInfo(chunks, param); CompressedInputStream input = new CompressedInputStream(new ByteArrayInputStream(toRead), info, ChecksumType.CRC32, () -> 1.0); - DataInputStream in = new DataInputStream(input); - for (int i = 0; i < sections.size(); i++) + try (DataInputStream in = new DataInputStream(input)) { - input.position(sections.get(i).left); - long readValue = in.readLong(); - assert readValue == valuesToCheck[i] : "expected " + valuesToCheck[i] + " but was " + readValue; + for (int i = 0; i < sections.size(); i++) + { + input.position(sections.get(i).left); + long readValue = in.readLong(); + assertEquals("expected " + valuesToCheck[i] + " but was " + readValue, valuesToCheck[i], readValue); + } } } }