diff --git a/bson-tests/src/commonMain/kotlin/raw/BinaryTest.kt b/bson-tests/src/commonMain/kotlin/raw/BinaryTest.kt index 3e3bbbbc..1a64cc3e 100644 --- a/bson-tests/src/commonMain/kotlin/raw/BinaryTest.kt +++ b/bson-tests/src/commonMain/kotlin/raw/BinaryTest.kt @@ -25,10 +25,7 @@ import opensavvy.ktmongo.bson.raw.BsonDeclaration.Companion.hex import opensavvy.ktmongo.bson.raw.BsonDeclaration.Companion.json import opensavvy.ktmongo.bson.raw.BsonDeclaration.Companion.serialize import opensavvy.ktmongo.bson.raw.BsonDeclaration.Companion.verify -import opensavvy.ktmongo.bson.types.ByteVector -import opensavvy.ktmongo.bson.types.FloatVector -import opensavvy.ktmongo.bson.types.UuidAsBsonBinarySerializer -import opensavvy.ktmongo.bson.types.Vector +import opensavvy.ktmongo.bson.types.* import opensavvy.ktmongo.dsl.LowLevelApi import opensavvy.prepared.suite.Prepared import opensavvy.prepared.suite.SuiteDsl @@ -334,6 +331,9 @@ fun SuiteDsl.binary(context: Prepared) = suite("Binary") { document { writeVector("x", Vector.fromBinaryData(Base64.decode("EAB/Bw=="))) }, + document { + writeVector("x", BooleanVector(true, true, true, true, true, true, true, false, true, true, true, false, false, false, false, false)) + }, hex("11000000057800040000000910007F0700"), json($$"""{"x": {"$binary": {"base64": "EAB/Bw==", "subType": "09"}}}"""), verify("Read type") { @@ -351,6 +351,9 @@ fun SuiteDsl.binary(context: Prepared) = suite("Binary") { verify("Read vector content") { check(read("x")?.readVector()?.raw.contentEquals(byteArrayOf(127, 7))) }, + verify("Read vector content as list") { + check(read("x")?.readVector() == listOf(true, true, true, true, true, true, true, false, true, true, true, false, false, false, false, false)) + }, ) testBson( @@ -430,6 +433,9 @@ fun SuiteDsl.binary(context: Prepared) = suite("Binary") { document { writeVector("x", Vector.fromBinaryData(Base64.decode("EAA="))) }, + document { + writeVector("x", BooleanVector()) + }, hex("0F0000000578000200000009100000"), json($$"""{"x": {"$binary": {"base64": "EAA=", "subType": "09"}}}"""), verify("Read type") { @@ -447,6 +453,9 @@ fun SuiteDsl.binary(context: Prepared) = suite("Binary") { verify("Read vector content") { check(read("x")?.readVector()?.raw?.size == 0) }, + verify("Read vector content as list") { + check(read("x")?.readVector() == emptyList()) + }, ) } diff --git a/bson/src/commonMain/kotlin/types/Vector.kt b/bson/src/commonMain/kotlin/types/Vector.kt index ddc69133..2a391b83 100644 --- a/bson/src/commonMain/kotlin/types/Vector.kt +++ b/bson/src/commonMain/kotlin/types/Vector.kt @@ -18,6 +18,8 @@ package opensavvy.ktmongo.bson.types import opensavvy.ktmongo.bson.BsonType import opensavvy.ktmongo.dsl.LowLevelApi +import kotlin.experimental.or +import kotlin.io.encoding.Base64 import kotlin.math.max import kotlin.math.min @@ -51,7 +53,7 @@ interface Vector { * Currently, the following types are implemented: * - `0x03`: [ByteVector] * - `0x27`: [FloatVector] - * - `0x10`: [PackedBitVector] + * - `0x10`: [BooleanVector] * * In most situations, users of this library should use `is` checks with one of the implementing subclasses * rather than attempting to match on this property. @@ -101,6 +103,7 @@ interface Vector { fun fromBinaryData(content: ByteArray): Vector = when (content[0]) { 0x03.toByte() -> ByteVector(content.sliceArray(2 until content.size), Unit) 0x27.toByte() -> FloatVector(content.sliceArray(2 until content.size)) + 0x10.toByte() -> BooleanVector(content.sliceArray(2 until content.size), content[1]) else -> UnknownVector(content) } } @@ -443,3 +446,190 @@ class FloatVector internal constructor( override fun toString(): String = joinToString(separator = ", ", prefix = "FloatVector[", postfix = "]") } + +private fun booleansToBytes(booleans: Collection): ByteArray { + val unpaddedSize = booleans.size / 8 + val hasPadding = booleans.size % 8 != 0 + val array = ByteArray(unpaddedSize + if (hasPadding) 1 else 0) + booleans.forEachIndexed { index, bool -> + val booleanIndex = index / 8 + val remainder = index % 8 + array[booleanIndex] = array[booleanIndex] or (if (bool) 1 shl remainder else 0).toByte() + } + return array +} + +/** + * A [Vector] of [Boolean] elements (BSON's `PackedBitVector`). + * + * The different bytes can be extracted with [toArray]. + * + * Alternatively, this class implements [List]. + */ +class BooleanVector internal constructor( + /** + * The underlying byte storage. **Do not mutate this array!** + * + * Note that this storage does NOT include the type nor the padding. + */ + private val rawUnsafe: ByteArray, + + @property:LowLevelApi + override val padding: Byte, +) : Vector, Iterable, Collection, List { + + init { + @OptIn(LowLevelApi::class) + require(padding in 0..7) { "A vector can only have a maximum padding of 1 byte (8 bits), but found: $padding declared bits" } + } + + constructor(booleans: Collection) : this( + rawUnsafe = booleansToBytes(booleans), + padding = (booleans.size % 8).toByte(), + ) + + constructor(vararg booleans: Boolean) : this(booleans.asList()) + + @LowLevelApi + override val type: Byte + get() = 0x10 + + @LowLevelApi + override val raw: ByteArray + get() = rawUnsafe.copyOf() + + override fun iterator(): BooleanIterator = + IteratorImpl() + + private inner class IteratorImpl : BooleanIterator() { + private var index = 0 + + override fun hasNext(): Boolean = + index < size + + override fun nextBoolean(): Boolean = + get(index++) + } + + @OptIn(LowLevelApi::class) + override val size: Int = + rawUnsafe.size * 8 - padding.toInt() + + override fun isEmpty(): Boolean = + size == 0 + + override fun contains(element: Boolean): Boolean { + for (i in indices) { + if (get(i) == element) + return true + } + return false + } + + override fun containsAll(elements: Collection): Boolean = + elements.all { contains(it) } + + override fun get(index: Int): Boolean { + if (index < 0) + throw IndexOutOfBoundsException("Index must be non-negative, found: $index") + + if (index >= size) + throw IndexOutOfBoundsException("Index must be less than size ($size), found: $index") + + val booleanIndex = index / 8 + val remainder = index % 8 + return rawUnsafe[booleanIndex].toInt() and (1 shl remainder) != 0 + } + + override fun indexOf(element: Boolean): Int { + for (i in indices) { + if (get(i) == element) + return i + } + return -1 + } + + override fun lastIndexOf(element: Boolean): Int { + for (i in size - 1 downTo 0) { + if (get(i) == element) return i + } + return -1 + } + + override fun listIterator(): ListIterator = + ListIteratorImpl() + + override fun listIterator(index: Int): ListIterator = + ListIteratorImpl(index) + + private inner class ListIteratorImpl( + private var index: Int = 0, + ) : ListIterator { + override fun next(): Boolean = + get(index++) + + override fun hasNext(): Boolean = + index < size + + override fun hasPrevious(): Boolean = + index > 0 + + override fun previous(): Boolean = + get(--index) + + override fun nextIndex(): Int = + min(index + 1, size) + + override fun previousIndex(): Int = + max(index - 1, 0) + } + + override fun subList(fromIndex: Int, toIndex: Int): List { + if (fromIndex < 0) + throw IndexOutOfBoundsException("fromIndex must be non-negative, found: $fromIndex") + + if (toIndex > size) + throw IndexOutOfBoundsException("toIndex must be less than size ($size), found: $toIndex") + + if (toIndex < fromIndex) + throw IllegalArgumentException("toIndex must be greater than or equal to fromIndex, found: toIndex=$toIndex, fromIndex=$fromIndex") + + val list = ArrayList(toIndex - fromIndex) + + for (i in fromIndex until toIndex) { + list += get(i) + } + + return list + } + + fun toArray(): BooleanArray = + BooleanArray(size) { get(it) } + + override fun equals(other: Any?): Boolean { + return when { + this === other -> true + other === null -> false + other is BooleanVector -> rawUnsafe.contentEquals(other.rawUnsafe) + other is List<*> -> { + if (size != other.size) return false + for (i in indices) { + if (get(i) != other[i]) return false + } + true + } + + else -> false + } + } + + override fun hashCode(): Int { + var hashCode = 1 + for (e in this) + hashCode = 31 * hashCode + e.hashCode() + return hashCode + } + + override fun toString(): String = + joinToString(separator = ", ", prefix = "BooleanVector[", postfix = "]") +}