mirror of
https://github.com/apple/pkl.git
synced 2026-08-26 21:54:03 +02:00
Enforce max collection sizes when decoding pkl-binary (#1831)
This commit is contained in:
@@ -41,16 +41,32 @@ public class PklBinaryDecoder extends AbstractPklBinaryDecoder {
|
|||||||
super(unpacker);
|
super(unpacker);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private PklBinaryDecoder(MessageUnpacker unpacker, int collectionSizeLimit) {
|
||||||
|
super(unpacker, collectionSizeLimit);
|
||||||
|
}
|
||||||
|
|
||||||
/** Decode a value from the supplied byte array. */
|
/** Decode a value from the supplied byte array. */
|
||||||
public static Object decode(byte[] bytes) {
|
public static Object decode(byte[] bytes) {
|
||||||
return new PklBinaryDecoder(MessagePack.newDefaultUnpacker(bytes)).decode();
|
return new PklBinaryDecoder(MessagePack.newDefaultUnpacker(bytes)).decode();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/** Decode a value from the supplied byte array. */
|
||||||
|
public static Object decode(byte[] bytes, int collectionSizeLimit) {
|
||||||
|
return new PklBinaryDecoder(MessagePack.newDefaultUnpacker(bytes), collectionSizeLimit)
|
||||||
|
.decode();
|
||||||
|
}
|
||||||
|
|
||||||
/** Decode a value from the supplied {@link InputStream}. */
|
/** Decode a value from the supplied {@link InputStream}. */
|
||||||
public static Object decode(InputStream inputStream) {
|
public static Object decode(InputStream inputStream) {
|
||||||
return new PklBinaryDecoder(MessagePack.newDefaultUnpacker(inputStream)).decode();
|
return new PklBinaryDecoder(MessagePack.newDefaultUnpacker(inputStream)).decode();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/** Decode a value from the supplied {@link InputStream}. */
|
||||||
|
public static Object decode(InputStream inputStream, int collectionSizeLimit) {
|
||||||
|
return new PklBinaryDecoder(MessagePack.newDefaultUnpacker(inputStream), collectionSizeLimit)
|
||||||
|
.decode();
|
||||||
|
}
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
protected RuntimeException doFail(Exception cause, long offset, List<String> path) {
|
protected RuntimeException doFail(Exception cause, long offset, List<String> path) {
|
||||||
return new RuntimeException(
|
return new RuntimeException(
|
||||||
|
|||||||
@@ -32,27 +32,32 @@ public final class CollectionUtils {
|
|||||||
|
|
||||||
@TruffleBoundary
|
@TruffleBoundary
|
||||||
public static <T> HashSet<T> newHashSet(int expectedSize) {
|
public static <T> HashSet<T> newHashSet(int expectedSize) {
|
||||||
return new HashSet<>((int) (expectedSize / LOAD_FACTOR) + 1, LOAD_FACTOR);
|
return new HashSet<>(calculateHashMapCapacity(expectedSize), LOAD_FACTOR);
|
||||||
}
|
}
|
||||||
|
|
||||||
@TruffleBoundary
|
@TruffleBoundary
|
||||||
public static <T> LinkedHashSet<T> newLinkedHashSet(int expectedSize) {
|
public static <T> LinkedHashSet<T> newLinkedHashSet(int expectedSize) {
|
||||||
return new LinkedHashSet<>((int) (expectedSize / LOAD_FACTOR) + 1, LOAD_FACTOR);
|
return new LinkedHashSet<>(calculateHashMapCapacity(expectedSize), LOAD_FACTOR);
|
||||||
}
|
}
|
||||||
|
|
||||||
@TruffleBoundary
|
@TruffleBoundary
|
||||||
public static <K, V> HashMap<K, V> newHashMap(int expectedSize) {
|
public static <K, V> HashMap<K, V> newHashMap(int expectedSize) {
|
||||||
return new HashMap<>((int) (expectedSize / LOAD_FACTOR) + 1, LOAD_FACTOR);
|
return new HashMap<>(calculateHashMapCapacity(expectedSize), LOAD_FACTOR);
|
||||||
}
|
}
|
||||||
|
|
||||||
@TruffleBoundary
|
@TruffleBoundary
|
||||||
public static <K extends @Nullable Object, V extends @Nullable Object>
|
public static <K extends @Nullable Object, V extends @Nullable Object>
|
||||||
LinkedHashMap<K, V> newLinkedHashMap(int expectedSize) {
|
LinkedHashMap<K, V> newLinkedHashMap(int expectedSize) {
|
||||||
return new LinkedHashMap<>((int) (expectedSize / LOAD_FACTOR) + 1, LOAD_FACTOR);
|
return new LinkedHashMap<>(calculateHashMapCapacity(expectedSize), LOAD_FACTOR);
|
||||||
}
|
}
|
||||||
|
|
||||||
@TruffleBoundary
|
@TruffleBoundary
|
||||||
public static <K, V> ConcurrentHashMap<K, V> newConcurrentHashMap(int expectedSize) {
|
public static <K, V> ConcurrentHashMap<K, V> newConcurrentHashMap(int expectedSize) {
|
||||||
return new ConcurrentHashMap<>((int) (expectedSize / LOAD_FACTOR) + 1, LOAD_FACTOR);
|
return new ConcurrentHashMap<>(calculateHashMapCapacity(expectedSize), LOAD_FACTOR);
|
||||||
|
}
|
||||||
|
|
||||||
|
private static int calculateHashMapCapacity(int expectedSize) {
|
||||||
|
// this is exactly how java.util.HashMap does it
|
||||||
|
return (int) Math.ceil(expectedSize / (double) LOAD_FACTOR);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
/*
|
/*
|
||||||
* Copyright © 2025 Apple Inc. and the Pkl project authors. All rights reserved.
|
* Copyright © 2025-2026 Apple Inc. and the Pkl project authors. All rights reserved.
|
||||||
*
|
*
|
||||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
* you may not use this file except in compliance with the License.
|
* you may not use this file except in compliance with the License.
|
||||||
@@ -40,10 +40,19 @@ import org.pkl.core.util.LateInit;
|
|||||||
*/
|
*/
|
||||||
public abstract class AbstractPklBinaryDecoder {
|
public abstract class AbstractPklBinaryDecoder {
|
||||||
private final MessageUnpacker unpacker;
|
private final MessageUnpacker unpacker;
|
||||||
|
private final int collectionSizeLimit;
|
||||||
@LateInit protected Deque<Object> currPath;
|
@LateInit protected Deque<Object> currPath;
|
||||||
|
|
||||||
|
private static final int defaultCollectionSizeLimit = 65535; // 16k
|
||||||
|
|
||||||
protected AbstractPklBinaryDecoder(MessageUnpacker unpacker) {
|
protected AbstractPklBinaryDecoder(MessageUnpacker unpacker) {
|
||||||
this.unpacker = unpacker;
|
this.unpacker = unpacker;
|
||||||
|
this.collectionSizeLimit = defaultCollectionSizeLimit;
|
||||||
|
}
|
||||||
|
|
||||||
|
protected AbstractPklBinaryDecoder(MessageUnpacker unpacker, int collectionSizeLimit) {
|
||||||
|
this.unpacker = unpacker;
|
||||||
|
this.collectionSizeLimit = collectionSizeLimit;
|
||||||
}
|
}
|
||||||
|
|
||||||
protected static class DecodeException extends RuntimeException {
|
protected static class DecodeException extends RuntimeException {
|
||||||
@@ -176,6 +185,13 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private void checkCollectionLength(int length, String collectionType) {
|
||||||
|
if (length <= collectionSizeLimit) return;
|
||||||
|
throw new DecodeException(
|
||||||
|
"Unable to decode %s of length %d, exceeded maximum collection size of %d",
|
||||||
|
collectionType, length, collectionSizeLimit);
|
||||||
|
}
|
||||||
|
|
||||||
private Object decodeObject(int len) throws IOException {
|
private Object decodeObject(int len) throws IOException {
|
||||||
assertLength(PklBinaryCode.OBJECT, len, 3);
|
assertLength(PklBinaryCode.OBJECT, len, 3);
|
||||||
currPath.push("'object");
|
currPath.push("'object");
|
||||||
@@ -201,7 +217,7 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
private Object decodeMap(int len) throws IOException {
|
private Object decodeMap(int len) throws IOException {
|
||||||
assertLength(PklBinaryCode.MAP, len, 1);
|
assertLength(PklBinaryCode.MAP, len, 1);
|
||||||
currPath.push("'map");
|
currPath.push("'map");
|
||||||
var result = doDecodeMap(new MapDecodeIterator(unpacker.unpackMapHeader()));
|
var result = doDecodeMap(new MapDecodeIterator(unpacker.unpackMapHeader(), "map"));
|
||||||
unpacker.skipValue(len - 2);
|
unpacker.skipValue(len - 2);
|
||||||
currPath.pop();
|
currPath.pop();
|
||||||
return result;
|
return result;
|
||||||
@@ -210,7 +226,7 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
private Object decodeMapping(int len) throws IOException {
|
private Object decodeMapping(int len) throws IOException {
|
||||||
assertLength(PklBinaryCode.MAPPING, len, 1);
|
assertLength(PklBinaryCode.MAPPING, len, 1);
|
||||||
currPath.push("'mapping");
|
currPath.push("'mapping");
|
||||||
var result = doDecodeMapping(new MapDecodeIterator(unpacker.unpackMapHeader()));
|
var result = doDecodeMapping(new MapDecodeIterator(unpacker.unpackMapHeader(), "mapping"));
|
||||||
unpacker.skipValue(len - 2);
|
unpacker.skipValue(len - 2);
|
||||||
currPath.pop();
|
currPath.pop();
|
||||||
return result;
|
return result;
|
||||||
@@ -219,7 +235,7 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
private Object decodeList(int len) throws IOException {
|
private Object decodeList(int len) throws IOException {
|
||||||
assertLength(PklBinaryCode.LIST, len, 1);
|
assertLength(PklBinaryCode.LIST, len, 1);
|
||||||
currPath.push("'list");
|
currPath.push("'list");
|
||||||
var result = doDecodeList(new CollectionDecodeIterator(unpacker.unpackArrayHeader()));
|
var result = doDecodeList(new CollectionDecodeIterator(unpacker.unpackArrayHeader(), "list"));
|
||||||
unpacker.skipValue(len - 2);
|
unpacker.skipValue(len - 2);
|
||||||
currPath.pop();
|
currPath.pop();
|
||||||
return result;
|
return result;
|
||||||
@@ -228,7 +244,8 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
private Object decodeListing(int len) throws IOException {
|
private Object decodeListing(int len) throws IOException {
|
||||||
assertLength(PklBinaryCode.LISTING, len, 1);
|
assertLength(PklBinaryCode.LISTING, len, 1);
|
||||||
currPath.push("'listing");
|
currPath.push("'listing");
|
||||||
var result = doDecodeListing(new CollectionDecodeIterator(unpacker.unpackArrayHeader()));
|
var result =
|
||||||
|
doDecodeListing(new CollectionDecodeIterator(unpacker.unpackArrayHeader(), "listing"));
|
||||||
unpacker.skipValue(len - 2);
|
unpacker.skipValue(len - 2);
|
||||||
currPath.pop();
|
currPath.pop();
|
||||||
return result;
|
return result;
|
||||||
@@ -237,7 +254,7 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
private Object decodeSet(int len) throws IOException {
|
private Object decodeSet(int len) throws IOException {
|
||||||
assertLength(PklBinaryCode.SET, len, 1);
|
assertLength(PklBinaryCode.SET, len, 1);
|
||||||
currPath.push("'set");
|
currPath.push("'set");
|
||||||
var result = doDecodeSet(new CollectionDecodeIterator(unpacker.unpackArrayHeader()));
|
var result = doDecodeSet(new CollectionDecodeIterator(unpacker.unpackArrayHeader(), "set"));
|
||||||
currPath.pop();
|
currPath.pop();
|
||||||
unpacker.skipValue(len - 2);
|
unpacker.skipValue(len - 2);
|
||||||
return result;
|
return result;
|
||||||
@@ -403,6 +420,7 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
protected class ObjectDecodeIterator extends DecodeIterator<DecodedObjectMember> {
|
protected class ObjectDecodeIterator extends DecodeIterator<DecodedObjectMember> {
|
||||||
ObjectDecodeIterator(int size) {
|
ObjectDecodeIterator(int size) {
|
||||||
super(size);
|
super(size);
|
||||||
|
checkCollectionLength(size, "object");
|
||||||
}
|
}
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
@@ -441,8 +459,9 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
protected class CollectionDecodeIterator extends DecodeIterator<Object> {
|
protected class CollectionDecodeIterator extends DecodeIterator<Object> {
|
||||||
CollectionDecodeIterator(int size) {
|
CollectionDecodeIterator(int size, String collectionType) {
|
||||||
super(size);
|
super(size);
|
||||||
|
checkCollectionLength(size, collectionType);
|
||||||
}
|
}
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
@@ -452,8 +471,9 @@ public abstract class AbstractPklBinaryDecoder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
protected class MapDecodeIterator extends DecodeIterator<Pair<Object, Object>> {
|
protected class MapDecodeIterator extends DecodeIterator<Pair<Object, Object>> {
|
||||||
MapDecodeIterator(int size) {
|
MapDecodeIterator(int size, String collectionType) {
|
||||||
super(size);
|
super(size);
|
||||||
|
checkCollectionLength(size, collectionType);
|
||||||
}
|
}
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
|
|||||||
@@ -230,4 +230,35 @@ class PklBinaryDecoderTest {
|
|||||||
"Unexpected blank typealias module URI",
|
"Unexpected blank typealias module URI",
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@Test
|
||||||
|
fun `decode collections too large`() {
|
||||||
|
assertExceptionCauseMessageContains(
|
||||||
|
byteArrayOf(0x93.toByte(), PklBinaryCode.OBJECT.code) +
|
||||||
|
strFoo +
|
||||||
|
strFoo +
|
||||||
|
byteArrayOf(0xdd.toByte(), 0, 1, 0, 0),
|
||||||
|
"Unable to decode object of length 65536, exceeded maximum collection size",
|
||||||
|
)
|
||||||
|
assertExceptionCauseMessageContains(
|
||||||
|
byteArrayOf(0x93.toByte(), PklBinaryCode.MAP.code, 0xdf.toByte(), 0, 1, 0, 0),
|
||||||
|
"Unable to decode map of length 65536, exceeded maximum collection size",
|
||||||
|
)
|
||||||
|
assertExceptionCauseMessageContains(
|
||||||
|
byteArrayOf(0x93.toByte(), PklBinaryCode.MAPPING.code, 0xdf.toByte(), 0, 1, 0, 0),
|
||||||
|
"Unable to decode mapping of length 65536, exceeded maximum collection size",
|
||||||
|
)
|
||||||
|
assertExceptionCauseMessageContains(
|
||||||
|
byteArrayOf(0x93.toByte(), PklBinaryCode.LIST.code, 0xdd.toByte(), 0, 1, 0, 0),
|
||||||
|
"Unable to decode list of length 65536, exceeded maximum collection size",
|
||||||
|
)
|
||||||
|
assertExceptionCauseMessageContains(
|
||||||
|
byteArrayOf(0x93.toByte(), PklBinaryCode.LISTING.code, 0xdd.toByte(), 0, 1, 0, 0),
|
||||||
|
"Unable to decode listing of length 65536, exceeded maximum collection size",
|
||||||
|
)
|
||||||
|
assertExceptionCauseMessageContains(
|
||||||
|
byteArrayOf(0x93.toByte(), PklBinaryCode.SET.code, 0xdd.toByte(), 0, 1, 0, 0),
|
||||||
|
"Unable to decode set of length 65536, exceeded maximum collection size",
|
||||||
|
)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user