Skip to content

Commit 9d969b9

Browse files
committed
add zstd codec for v2
1 parent d8a7c92 commit 9d969b9

6 files changed

Lines changed: 84 additions & 29 deletions

File tree

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,35 @@
1+
package dev.zarr.zarrjava.core.codec.core;
2+
3+
import com.github.luben.zstd.Zstd;
4+
import com.github.luben.zstd.ZstdCompressCtx;
5+
import dev.zarr.zarrjava.ZarrException;
6+
import dev.zarr.zarrjava.core.codec.BytesBytesCodec;
7+
import dev.zarr.zarrjava.utils.Utils;
8+
9+
import java.nio.ByteBuffer;
10+
11+
public abstract class ZstdCodec extends BytesBytesCodec {
12+
13+
@Override
14+
public ByteBuffer decode(ByteBuffer compressedBytes) throws ZarrException {
15+
byte[] compressedArray = Utils.toArray(compressedBytes);
16+
long originalSize = Zstd.getFrameContentSize(compressedArray);
17+
if (originalSize < 0) {
18+
throw new ZarrException("Failed to get decompressed zstd size.");
19+
}
20+
byte[] decompressed = Zstd.decompress(compressedArray, (int) originalSize);
21+
return ByteBuffer.wrap(decompressed);
22+
}
23+
24+
protected ByteBuffer encodeInternal(int level, boolean checksum, ByteBuffer chunkBytes)
25+
throws ZarrException {
26+
byte[] arr = Utils.toArray(chunkBytes);
27+
byte[] compressed;
28+
try (ZstdCompressCtx ctx = new ZstdCompressCtx()) {
29+
ctx.setLevel(level);
30+
ctx.setChecksum(checksum);
31+
compressed = ctx.compress(arr);
32+
}
33+
return ByteBuffer.wrap(compressed);
34+
}
35+
}

src/main/java/dev/zarr/zarrjava/v2/codec/CodecRegistry.java

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import com.fasterxml.jackson.databind.jsontype.NamedType;
44
import dev.zarr.zarrjava.v2.codec.core.BloscCodec;
55
import dev.zarr.zarrjava.v2.codec.core.ZlibCodec;
6+
import dev.zarr.zarrjava.v2.codec.core.ZstdCodec;
67

78
import java.util.HashMap;
89
import java.util.Map;
@@ -14,6 +15,7 @@ public class CodecRegistry {
1415
static {
1516
addType("blosc", BloscCodec.class);
1617
addType("zlib", ZlibCodec.class);
18+
addType("zstd", ZstdCodec.class);
1719
}
1820

1921
public static void addType(String name, Class<? extends Codec> codecClass) {
Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
package dev.zarr.zarrjava.v2.codec.core;
2+
3+
import com.fasterxml.jackson.annotation.JsonCreator;
4+
import com.fasterxml.jackson.annotation.JsonIgnore;
5+
import com.fasterxml.jackson.annotation.JsonProperty;
6+
import dev.zarr.zarrjava.ZarrException;
7+
import dev.zarr.zarrjava.core.ArrayMetadata;
8+
import dev.zarr.zarrjava.v2.codec.Codec;
9+
10+
import java.nio.ByteBuffer;
11+
12+
public class ZstdCodec extends dev.zarr.zarrjava.core.codec.core.ZstdCodec implements Codec {
13+
14+
@JsonIgnore
15+
public final String id = "zstd";
16+
public final int level;
17+
public final boolean checksum;
18+
19+
@JsonCreator(mode = JsonCreator.Mode.PROPERTIES)
20+
public ZstdCodec(
21+
@JsonProperty(value = "level", defaultValue = "0") int level,
22+
@JsonProperty(value = "checksum", defaultValue = "false") boolean checksum) throws ZarrException {
23+
if (level < -131072 || level > 22) {
24+
throw new ZarrException("'level' needs to be between -131072 and 22.");
25+
}
26+
this.level = level;
27+
this.checksum = checksum;
28+
}
29+
30+
@Override
31+
public ByteBuffer encode(ByteBuffer chunkBytes) throws ZarrException {
32+
return encodeInternal(this.level, this.checksum, chunkBytes);
33+
}
34+
35+
@Override
36+
public Codec evolveFromCoreArrayMetadata(ArrayMetadata.CoreArrayMetadata arrayMetadata) {
37+
return this;
38+
}
39+
}

src/main/java/dev/zarr/zarrjava/v3/codec/core/ZstdCodec.java

Lines changed: 2 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -3,18 +3,14 @@
33
import com.fasterxml.jackson.annotation.JsonCreator;
44
import com.fasterxml.jackson.annotation.JsonIgnore;
55
import com.fasterxml.jackson.annotation.JsonProperty;
6-
import com.github.luben.zstd.Zstd;
7-
import com.github.luben.zstd.ZstdCompressCtx;
86
import dev.zarr.zarrjava.ZarrException;
9-
import dev.zarr.zarrjava.core.codec.BytesBytesCodec;
10-
import dev.zarr.zarrjava.utils.Utils;
117
import dev.zarr.zarrjava.v3.ArrayMetadata;
128
import dev.zarr.zarrjava.v3.codec.Codec;
139

1410
import javax.annotation.Nonnull;
1511
import java.nio.ByteBuffer;
1612

17-
public class ZstdCodec extends BytesBytesCodec implements Codec {
13+
public class ZstdCodec extends dev.zarr.zarrjava.core.codec.core.ZstdCodec implements Codec {
1814

1915
@JsonIgnore
2016
public final String name = "zstd";
@@ -27,29 +23,9 @@ public ZstdCodec(
2723
this.configuration = configuration;
2824
}
2925

30-
@Override
31-
public ByteBuffer decode(ByteBuffer compressedBytes) throws ZarrException {
32-
byte[] compressedArray = Utils.toArray(compressedBytes);
33-
34-
long originalSize = Zstd.getFrameContentSize(compressedArray);
35-
if (originalSize == 0) {
36-
throw new ZarrException("Failed to get decompressed size");
37-
}
38-
39-
byte[] decompressed = Zstd.decompress(compressedArray, (int) originalSize);
40-
return ByteBuffer.wrap(decompressed);
41-
}
42-
4326
@Override
4427
public ByteBuffer encode(ByteBuffer chunkBytes) throws ZarrException {
45-
byte[] arr = Utils.toArray(chunkBytes);
46-
byte[] compressed;
47-
try (ZstdCompressCtx ctx = new ZstdCompressCtx()) {
48-
ctx.setLevel(configuration.level);
49-
ctx.setChecksum(configuration.checksum);
50-
compressed = ctx.compress(arr);
51-
}
52-
return ByteBuffer.wrap(compressed);
28+
return encodeInternal(configuration.level, configuration.checksum, chunkBytes);
5329
}
5430

5531
@Override
@@ -75,5 +51,3 @@ public Configuration(@JsonProperty(value = "level", defaultValue = "5") int leve
7551
}
7652
}
7753
}
78-
79-

src/test/java/dev/zarr/zarrjava/ZarrPythonTests.java

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -106,7 +106,9 @@ static Stream<Object[]> compressorAndDataTypeProviderV2() {
106106
new Object[]{"blosc", "lz4_shuffle_6", dev.zarr.zarrjava.v2.DataType.INT32},
107107
new Object[]{"blosc", "lz4hc_bitshuffle_3", dev.zarr.zarrjava.v2.DataType.INT32},
108108
new Object[]{"blosc", "zlib_shuffle_5", dev.zarr.zarrjava.v2.DataType.INT32},
109-
new Object[]{"blosc", "zstd_bitshuffle_9", dev.zarr.zarrjava.v2.DataType.INT32}
109+
new Object[]{"blosc", "zstd_bitshuffle_9", dev.zarr.zarrjava.v2.DataType.INT32},
110+
new Object[]{"zstd", "0_true", dev.zarr.zarrjava.v2.DataType.INT32},
111+
new Object[]{"zstd", "5_false", dev.zarr.zarrjava.v2.DataType.INT32}
110112
);
111113

112114
return Stream.concat(datatypeTests, bloscTests);

src/test/python-scripts/parse_codecs.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,9 @@ def parse_codecs_zarr_python(codec_string: str, param_string: str, zarr_version:
4848
codecs=[BytesCodec(endian="little")]),))
4949
elif codec_string == "crc32c" and zarr_version == 3:
5050
compressor = Crc32cCodec()
51+
elif codec_string == "zstd" and zarr_version == 2:
52+
level, checksum = param_string.split("_")
53+
compressor = numcodecs.Zstd(level=int(level), checksum=checksum == 'true')
5154
else:
5255
raise ValueError(f"Invalid codec: {codec_string}, zarr_version: {zarr_version}")
5356

0 commit comments

Comments
 (0)