diff --git a/benchmark/src/androidTest/java/com/vitorpamplona/amethyst/benchmark/ThumbHashBenchmark.kt b/benchmark/src/androidTest/java/com/vitorpamplona/amethyst/benchmark/ThumbHashBenchmark.kt new file mode 100644 index 000000000..e7aaa4d56 --- /dev/null +++ b/benchmark/src/androidTest/java/com/vitorpamplona/amethyst/benchmark/ThumbHashBenchmark.kt @@ -0,0 +1,156 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.vitorpamplona.amethyst.benchmark + +import androidx.benchmark.junit4.BenchmarkRule +import androidx.benchmark.junit4.measureRepeated +import androidx.test.ext.junit.runners.AndroidJUnit4 +import com.vitorpamplona.amethyst.commons.thumbhash.ThumbHashDecoder +import com.vitorpamplona.amethyst.commons.thumbhash.ThumbHashEncoder +import org.junit.Rule +import org.junit.Test +import org.junit.runner.RunWith +import kotlin.io.encoding.Base64 +import kotlin.io.encoding.ExperimentalEncodingApi + +@OptIn(ExperimentalEncodingApi::class) +@RunWith(AndroidJUnit4::class) +class ThumbHashBenchmark { + @get:Rule + val benchmarkRule = BenchmarkRule() + + // Representative opaque landscape hash. Produced by encoding a smooth + // 32x24 warm gradient so the AC coefficients exercise the full L block. + private val warmLandscape = + run { + val w = 32 + val h = 24 + val pixels = + IntArray(w * h) { i -> + val x = i % w + val y = i / w + val r = (160 + (x * 3)) and 0xff + val g = (90 + (y * 4)) and 0xff + val b = (40 + ((x + y) * 2)) and 0xff + (0xFF shl 24) or (r shl 16) or (g shl 8) or b + } + ThumbHashEncoder.encodeToBase64(pixels, w, h) + } + + // Representative alpha hash. Radial alpha vignette forces the alpha DCT + // block to carry real energy so decode cost matches production inputs. + private val alphaPortrait = + run { + val w = 24 + val h = 32 + val pixels = + IntArray(w * h) { i -> + val x = i % w + val y = i / w + val dx = x - w / 2 + val dy = y - h / 2 + val dist = kotlin.math.sqrt((dx * dx + dy * dy).toDouble()) + val maxDist = kotlin.math.sqrt((w * w / 4 + h * h / 4).toDouble()) + val alpha = (255 * (1.0 - (dist / maxDist).coerceIn(0.0, 1.0))).toInt() + (alpha shl 24) or (0x40 shl 16) or (0x80 shl 8) or 0xC0 + } + ThumbHashEncoder.encodeToBase64(pixels, w, h) + } + + private val warmBytes = Base64.decode(padded(warmLandscape)) + private val alphaBytes = Base64.decode(padded(alphaPortrait)) + + private fun padded(s: String): String { + val r = s.length % 4 + return if (r == 0) s else s + "=".repeat(4 - r) + } + + @Test + fun testAspectRatioFromBase64() { + // Warm up + ThumbHashDecoder.aspectRatio(warmLandscape) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.aspectRatio(warmLandscape) + } + } + + @Test + fun testAspectRatioFromBytes() { + ThumbHashDecoder.aspectRatio(warmBytes) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.aspectRatio(warmBytes) + } + } + + @Test + fun testDecodeOpaqueBytes() { + // Warm up the cosine cache for this size. + ThumbHashDecoder.decode(warmBytes) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.decode(warmBytes) + } + } + + @Test + fun testDecodeWithAlphaBytes() { + ThumbHashDecoder.decode(alphaBytes) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.decode(alphaBytes) + } + } + + @Test + fun testDecodeOpaqueBase64() { + ThumbHashDecoder.decode(warmLandscape) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.decode(warmLandscape) + } + } + + /** + * Measures decode cost when the cosine cache is cold on every call. + * This is the realistic "first time we see a new output size" cost; the + * cached case above represents steady-state feed scrolling. + */ + @Test + fun testDecodeOpaqueColdCache() { + ThumbHashDecoder.decode(warmBytes) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.clearCache() + ThumbHashDecoder.decode(warmBytes) + } + } + + @Test + fun testDecodeKeepAspectRatio() { + ThumbHashDecoder.decodeKeepAspectRatio(warmLandscape, 32) + + benchmarkRule.measureRepeated { + ThumbHashDecoder.decodeKeepAspectRatio(warmLandscape, 32) + } + } +} diff --git a/commons/src/commonMain/kotlin/com/vitorpamplona/amethyst/commons/thumbhash/ThumbHashDecoder.kt b/commons/src/commonMain/kotlin/com/vitorpamplona/amethyst/commons/thumbhash/ThumbHashDecoder.kt index 5cde68231..30bc6fd3b 100644 --- a/commons/src/commonMain/kotlin/com/vitorpamplona/amethyst/commons/thumbhash/ThumbHashDecoder.kt +++ b/commons/src/commonMain/kotlin/com/vitorpamplona/amethyst/commons/thumbhash/ThumbHashDecoder.kt @@ -26,16 +26,63 @@ import kotlin.io.encoding.ExperimentalEncodingApi import kotlin.math.PI import kotlin.math.cos import kotlin.math.max -import kotlin.math.min import kotlin.math.round /** * ThumbHash decoder. * * Port of the reference implementation by Evan Wallace - * (https://github.com/evanw/thumbhash, public domain), adapted to Kotlin. + * (https://github.com/evanw/thumbhash, public domain), with performance + * optimisations for the decode hot path: + * + * - cosine tables for the inverse DCT are precomputed once per decode and + * cached across decodes keyed by `(size, componentCount)`; the reference + * JS impl recomputes them for every single output pixel. + * - AC coefficients are unpacked into fixed-size `DoubleArray`s, avoiding + * `ArrayList` boxing and the final array copy. + * - The LPQA → sRGB conversion uses an inline branch clamp instead of + * `min/max/round/coerceIn` chains. */ object ThumbHashDecoder { + // Cosine tables are small and decoded sizes repeat heavily in practice + // (every Coil request at a given target width shares the same table). + // Keep an unbounded map — there are at most a few dozen distinct + // (size, components) pairs across the entire app lifetime, each table is + // a few KB, so the memory ceiling is tiny. + private val cosineCache = HashMap() + private val cosineCacheLock = Any() + + /** + * Clear the cosine table cache. Tables are tiny but callers under memory + * pressure can release them; they will be recomputed on demand. + */ + fun clearCache() { + synchronized(cosineCacheLock) { cosineCache.clear() } + } + + private fun cosTable( + size: Int, + components: Int, + ): DoubleArray { + val key = (size.toLong() shl 32) or components.toLong() + synchronized(cosineCacheLock) { + cosineCache[key]?.let { return it } + } + val table = DoubleArray(size * components) + val piOverSize = PI / size + for (i in 0 until size) { + val phase = piOverSize * (i + 0.5) + val rowOffset = i * components + for (c in 0 until components) { + table[rowOffset + c] = cos(phase * c) + } + } + synchronized(cosineCacheLock) { + cosineCache.getOrPut(key) { table } + } + return table + } + /** * Returns width/height. Returns null if the hash is malformed. */ @@ -67,7 +114,20 @@ object ThumbHashDecoder { val width: Int, val height: Int, val pixels: IntArray, - ) + ) { + override fun equals(other: Any?): Boolean { + if (this === other) return true + if (other !is RGBAImage) return false + return width == other.width && height == other.height && pixels.contentEquals(other.pixels) + } + + override fun hashCode(): Int { + var result = width + result = 31 * result + height + result = 31 * result + pixels.contentHashCode() + return result + } + } /** * Decode a ThumbHash byte array into ARGB pixels. @@ -106,100 +166,106 @@ object ThumbHashDecoder { aScale = 0.0 } + // Pre-size and unpack AC coefficients + val lAcCount = countAc(lx, ly) + val pqAcCount = countAc(3, 3) + val aAcCount = if (hasAlpha) countAc(5, 5) else 0 + val totalAc = lAcCount + pqAcCount * 2 + aAcCount val acStart = if (hasAlpha) 6 else 5 + val acBytesAvailable = hash.size - acStart + // 2 coefficients per byte + if (acBytesAvailable * 2 < totalAc) return null + + val lAc = DoubleArray(lAcCount) + val pAc = DoubleArray(pqAcCount) + val qAc = DoubleArray(pqAcCount) + val aAc = if (hasAlpha) DoubleArray(aAcCount) else EMPTY_DOUBLE + var acIndex = 0 + acIndex = readAcInto(hash, acStart, acIndex, lx, ly, lScale, lAc) + acIndex = readAcInto(hash, acStart, acIndex, 3, 3, pScale * 1.25, pAc) + acIndex = readAcInto(hash, acStart, acIndex, 3, 3, qScale * 1.25, qAc) + if (hasAlpha) readAcInto(hash, acStart, acIndex, 5, 5, aScale, aAc) - fun readAc( - nx: Int, - ny: Int, - scale: Double, - ): DoubleArray { - val ac = ArrayList(nx * ny) - var cy = 0 - while (cy < ny) { - var cx = if (cy != 0) 0 else 1 - while (cx * ny < nx * (ny - cy)) { - val byteIdx = acStart + (acIndex shr 1) - if (byteIdx >= hash.size) return DoubleArray(0) - val shift = (acIndex and 1) shl 2 - val q4 = ((hash[byteIdx].toInt() ushr shift) and 15) - ac.add((q4 / 7.5 - 1.0) * scale) - acIndex++ - cx++ - } - cy++ - } - return DoubleArray(ac.size) { ac[it] } - } - - val lAc = readAc(lx, ly, lScale) - val pAc = readAc(3, 3, pScale * 1.25) - val qAc = readAc(3, 3, qScale * 1.25) - val aAc = if (hasAlpha) readAc(5, 5, aScale) else DoubleArray(0) - + // Output size val ratio = lx.toDouble() / ly.toDouble() val w = round(if (ratio > 1) 32.0 else 32.0 * ratio).toInt() val h = round(if (ratio > 1) 32.0 / ratio else 32.0).toInt() val pixels = IntArray(w * h) - val fxMax = max(lx, if (hasAlpha) 5 else 3) - val fyMax = max(ly, if (hasAlpha) 5 else 3) - val fx = DoubleArray(fxMax) - val fy = DoubleArray(fyMax) + // Precomputed cosine tables (shared across decodes with matching size/components) + val cosXL = cosTable(w, lx) + val cosYL = cosTable(h, ly) + val cosXPQ = cosTable(w, 3) + val cosYPQ = cosTable(h, 3) + val cosXA: DoubleArray + val cosYA: DoubleArray + if (hasAlpha) { + cosXA = cosTable(w, 5) + cosYA = cosTable(h, 5) + } else { + cosXA = EMPTY_DOUBLE + cosYA = EMPTY_DOUBLE + } + // Decode pixels using the inverse DCT + var pixelIdx = 0 for (y in 0 until h) { + val cosYLBase = y * ly + val cosYPQBase = y * 3 + val cosYABase = y * 5 for (x in 0 until w) { - var lVal = lDc - var pVal = pDc - var qVal = qDc - var aVal = aDc + val cosXLBase = x * lx + val cosXPQBase = x * 3 + val cosXABase = x * 5 - for (cx in 0 until fxMax) fx[cx] = cos(PI / w * (x + 0.5) * cx) - for (cy in 0 until fyMax) fy[cy] = cos(PI / h * (y + 0.5) * cy) + var l = lDc + var p = pDc + var q = qDc + var a = aDc - // L - run { - var cy = 0 - var j = 0 - while (cy < ly) { - var cx = if (cy != 0) 0 else 1 - val fy2 = fy[cy] * 2.0 - while (cx * ly < lx * (ly - cy)) { - lVal += lAc[j] * fx[cx] * fy2 - j++ - cx++ - } - cy++ + // L channel — triangular iteration over (cx, cy) + var j = 0 + var cy = 0 + while (cy < ly) { + val fyL2 = cosYL[cosYLBase + cy] * 2.0 + var cx = if (cy != 0) 0 else 1 + val cxLimit = cxLimitForL(lx, ly, cy) + while (cx < cxLimit) { + l += lAc[j] * cosXL[cosXLBase + cx] * fyL2 + j++ + cx++ } + cy++ } - // P & Q - run { - var cy = 0 - var j = 0 - while (cy < 3) { - var cx = if (cy != 0) 0 else 1 - val fy2 = fy[cy] * 2.0 - while (cx < 3 - cy) { - val f = fx[cx] * fy2 - pVal += pAc[j] * f - qVal += qAc[j] * f - j++ - cx++ - } - cy++ + // P and Q share the same 3x3 triangular iteration + j = 0 + cy = 0 + while (cy < 3) { + val fyPQ2 = cosYPQ[cosYPQBase + cy] * 2.0 + var cx = if (cy != 0) 0 else 1 + val cxLimit = 3 - cy + while (cx < cxLimit) { + val f = cosXPQ[cosXPQBase + cx] * fyPQ2 + p += pAc[j] * f + q += qAc[j] * f + j++ + cx++ } + cy++ } - // A + // Alpha channel if (hasAlpha) { - var cy = 0 - var j = 0 + j = 0 + cy = 0 while (cy < 5) { + val fyA2 = cosYA[cosYABase + cy] * 2.0 var cx = if (cy != 0) 0 else 1 - val fy2 = fy[cy] * 2.0 - while (cx < 5 - cy) { - aVal += aAc[j] * fx[cx] * fy2 + val cxLimit = 5 - cy + while (cx < cxLimit) { + a += aAc[j] * cosXA[cosXABase + cx] * fyA2 j++ cx++ } @@ -207,15 +273,15 @@ object ThumbHashDecoder { } } - // LPQA → RGB - val bCh = lVal - 2.0 / 3.0 * pVal - val rCh = (3.0 * lVal - bCh + qVal) / 2.0 - val gCh = rCh - qVal - val rOut = (255.0 * min(1.0, max(0.0, rCh))).let { round(it).toInt() }.coerceIn(0, 255) - val gOut = (255.0 * min(1.0, max(0.0, gCh))).let { round(it).toInt() }.coerceIn(0, 255) - val bOut = (255.0 * min(1.0, max(0.0, bCh))).let { round(it).toInt() }.coerceIn(0, 255) - val aOut = if (hasAlpha) (255.0 * min(1.0, max(0.0, aVal))).let { round(it).toInt() }.coerceIn(0, 255) else 255 - pixels[x + y * w] = (aOut shl 24) or (rOut shl 16) or (gOut shl 8) or bOut + // LPQA → sRGB with inline clamp + val bCh = l - 2.0 / 3.0 * p + val rCh = (3.0 * l - bCh + q) * 0.5 + val gCh = rCh - q + val rOut = clamp255(rCh) + val gOut = clamp255(gCh) + val bOut = clamp255(bCh) + val aOut = if (hasAlpha) clamp255(a) else 255 + pixels[pixelIdx++] = (aOut shl 24) or (rOut shl 16) or (gOut shl 8) or bOut } } @@ -236,12 +302,13 @@ object ThumbHashDecoder { } /** - * Decode a ThumbHash string into a [PlatformImage] roughly [targetWidth] wide, - * preserving the aspect ratio of the original image. - * - * Mirrors [com.vitorpamplona.amethyst.commons.blurhash.BlurHashDecoder.decodeKeepAspectRatio] - * so existing placeholder pipelines can swap in thumbhash transparently. + * Decode a ThumbHash string into a [PlatformImage] whose aspect ratio + * matches the original image. [targetWidth] is accepted for API symmetry + * with [com.vitorpamplona.amethyst.commons.blurhash.BlurHashDecoder.decodeKeepAspectRatio] + * but the intrinsic decode output size is used because ThumbHash's own + * reconstruction is already aspect-correct at ~32px. */ + @Suppress("UNUSED_PARAMETER") fun decodeKeepAspectRatio( hash: String?, targetWidth: Int, @@ -250,6 +317,88 @@ object ThumbHashDecoder { return PlatformImage.create(rgba.pixels, rgba.width, rgba.height) } + // --- internal helpers --- // + + private val EMPTY_DOUBLE = DoubleArray(0) + + /** + * Count the number of AC coefficients carried by a channel of size nx × ny, + * following the reference implementation's triangular traversal. + */ + private fun countAc( + nx: Int, + ny: Int, + ): Int { + var count = 0 + var cy = 0 + while (cy < ny) { + var cx = if (cy != 0) 0 else 1 + while (cx * ny < nx * (ny - cy)) { + count++ + cx++ + } + cy++ + } + return count + } + + /** + * Row limit for the L channel's triangular traversal. For nx == ny this + * collapses to `nx - cy`; keeping the explicit form avoids a mispredicted + * branch in the inner loop for non-square L blocks. + */ + private fun cxLimitForL( + lx: Int, + ly: Int, + cy: Int, + ): Int { + // cx * ly < lx * (ly - cy) ⇔ cx < (lx * (ly - cy)) / ly + // Use integer ceil emulation: smallest cx that fails the condition. + val numerator = lx * (ly - cy) + // Largest cx satisfying cx * ly < numerator: + // cx <= ceil(numerator / ly) - 1 when numerator is exact, + // otherwise cx <= floor(numerator / ly). + // So the limit (exclusive) is ceil(numerator / ly) when numerator % ly != 0, + // else numerator / ly. + return if (numerator % ly == 0) numerator / ly else numerator / ly + 1 + } + + private fun readAcInto( + hash: ByteArray, + acStart: Int, + startIndex: Int, + nx: Int, + ny: Int, + scale: Double, + out: DoubleArray, + ): Int { + var acIndex = startIndex + var outIdx = 0 + val hashLen = hash.size + var cy = 0 + while (cy < ny) { + var cx = if (cy != 0) 0 else 1 + while (cx * ny < nx * (ny - cy)) { + val byteIdx = acStart + (acIndex shr 1) + if (byteIdx >= hashLen) return acIndex + val shift = (acIndex and 1) shl 2 + val q4 = (hash[byteIdx].toInt() ushr shift) and 15 + out[outIdx++] = (q4 / 7.5 - 1.0) * scale + acIndex++ + cx++ + } + cy++ + } + return acIndex + } + + /** Clamp v into 0..1 and scale to 0..255 with rounding, branchlessly on the hot path. */ + private fun clamp255(v: Double): Int { + if (v <= 0.0) return 0 + if (v >= 1.0) return 255 + return (v * 255.0 + 0.5).toInt() + } + private fun padBase64(s: String): String { val remainder = s.length % 4 return if (remainder == 0) s else s + "=".repeat(4 - remainder) diff --git a/commons/src/commonTest/kotlin/com/vitorpamplona/amethyst/commons/ThumbHashTest.kt b/commons/src/commonTest/kotlin/com/vitorpamplona/amethyst/commons/ThumbHashTest.kt index d374f0b5d..a152963ec 100644 --- a/commons/src/commonTest/kotlin/com/vitorpamplona/amethyst/commons/ThumbHashTest.kt +++ b/commons/src/commonTest/kotlin/com/vitorpamplona/amethyst/commons/ThumbHashTest.kt @@ -24,8 +24,10 @@ import com.vitorpamplona.amethyst.commons.thumbhash.ThumbHashDecoder import com.vitorpamplona.amethyst.commons.thumbhash.ThumbHashEncoder import org.junit.Assert.assertEquals import org.junit.Assert.assertNotNull +import org.junit.Assert.assertNull import org.junit.Assert.assertTrue import org.junit.Test +import kotlin.math.abs class ThumbHashTest { @Test @@ -84,8 +86,166 @@ class ThumbHashTest { @Test fun `decoding a malformed hash returns null`() { - assertEquals(null, ThumbHashDecoder.decode(ByteArray(3))) - assertEquals(null, ThumbHashDecoder.decode(null as String?)) - assertEquals(null, ThumbHashDecoder.decode("")) + assertNull(ThumbHashDecoder.decode(ByteArray(3))) + assertNull(ThumbHashDecoder.decode(null as String?)) + assertNull(ThumbHashDecoder.decode("")) + } + + @Test + fun `decoded opaque image has fully opaque alpha`() { + val w = 32 + val h = 32 + val pixels = IntArray(w * h) { 0xFF8080FF.toInt() } // opaque cornflower-ish + val decoded = ThumbHashDecoder.decode(ThumbHashEncoder.encode(pixels, w, h)) + assertNotNull(decoded) + decoded!! + for (p in decoded.pixels) { + val a = (p ushr 24) and 0xff + assertEquals("alpha should be 255 for opaque encode", 255, a) + } + } + + @Test + fun `decoded transparent image preserves alpha channel`() { + val w = 32 + val h = 32 + // Fully transparent pixels everywhere. + val pixels = IntArray(w * h) { 0x00000000 } + val decoded = ThumbHashDecoder.decode(ThumbHashEncoder.encode(pixels, w, h)) + assertNotNull(decoded) + decoded!! + // The average alpha is 0, so every decoded alpha should be at or near 0. + var maxAlpha = 0 + for (p in decoded.pixels) { + val a = (p ushr 24) and 0xff + if (a > maxAlpha) maxAlpha = a + } + assertTrue("max alpha of all-transparent decode should be low; got $maxAlpha", maxAlpha <= 16) + } + + @Test + fun `decoded average color is close to input average`() { + val w = 48 + val h = 32 + val target = intArrayOf(200, 120, 60) // warm orange + val pixels = + IntArray(w * h) { + (0xFF shl 24) or (target[0] shl 16) or (target[1] shl 8) or target[2] + } + val decoded = ThumbHashDecoder.decode(ThumbHashEncoder.encode(pixels, w, h)) + assertNotNull(decoded) + decoded!! + + var sumR = 0 + var sumG = 0 + var sumB = 0 + for (p in decoded.pixels) { + sumR += (p shr 16) and 0xff + sumG += (p shr 8) and 0xff + sumB += p and 0xff + } + val count = decoded.pixels.size + val avgR = sumR / count + val avgG = sumG / count + val avgB = sumB / count + + // ThumbHash quantisation allows a handful of codepoints of drift. + assertTrue("avg R drift: expected ${target[0]}, got $avgR", abs(avgR - target[0]) < 8) + assertTrue("avg G drift: expected ${target[1]}, got $avgG", abs(avgG - target[1]) < 8) + assertTrue("avg B drift: expected ${target[2]}, got $avgB", abs(avgB - target[2]) < 8) + } + + @Test + fun `aspect ratio matches landscape input`() { + val w = 60 + val h = 30 + val pixels = IntArray(w * h) { 0xFF446688.toInt() } + val hash = ThumbHashEncoder.encode(pixels, w, h) + val ratio = ThumbHashDecoder.aspectRatio(hash) + assertNotNull(ratio) + assertTrue("landscape ratio should be > 1, got $ratio", ratio!! > 1f) + } + + @Test + fun `aspect ratio matches portrait input`() { + val w = 30 + val h = 60 + val pixels = IntArray(w * h) { 0xFF446688.toInt() } + val hash = ThumbHashEncoder.encode(pixels, w, h) + val ratio = ThumbHashDecoder.aspectRatio(hash) + assertNotNull(ratio) + assertTrue("portrait ratio should be < 1, got $ratio", ratio!! < 1f) + } + + @Test + fun `repeated decodes produce identical output (cosine cache determinism)`() { + val w = 40 + val h = 30 + val pixels = + IntArray(w * h) { i -> + val x = i % w + (0xFF shl 24) or (x * 6 shl 16) or ((i % 255) shl 8) or ((i * 3) and 0xff) + } + val hash = ThumbHashEncoder.encode(pixels, w, h) + + val first = ThumbHashDecoder.decode(hash) + val second = ThumbHashDecoder.decode(hash) + val third = ThumbHashDecoder.decode(hash) + assertNotNull(first) + assertNotNull(second) + assertNotNull(third) + + // Bit-exact: the cached cosine tables must produce identical output. + assertEquals(first, second) + assertEquals(first, third) + } + + @Test + fun `clearCache does not affect correctness of subsequent decodes`() { + val w = 32 + val h = 32 + val pixels = IntArray(w * h) { 0xFFAABBCC.toInt() } + val hash = ThumbHashEncoder.encode(pixels, w, h) + + val before = ThumbHashDecoder.decode(hash) + ThumbHashDecoder.clearCache() + val after = ThumbHashDecoder.decode(hash) + assertEquals(before, after) + } + + @Test + fun `truncated AC payload returns null`() { + val w = 32 + val h = 32 + val pixels = IntArray(w * h) { 0xFF336699.toInt() } + val fullHash = ThumbHashEncoder.encode(pixels, w, h) + // Chop off half the AC payload. + val truncated = fullHash.copyOfRange(0, 5 + (fullHash.size - 5) / 4) + assertNull( + "hash with insufficient AC bytes should be rejected", + ThumbHashDecoder.decode(truncated), + ) + } + + @Test + fun `decoded output size stays within 32 x 32 bounds`() { + val w = 50 + val h = 40 + val pixels = + IntArray(w * h) { i -> + (0xFF shl 24) or ((i and 0xff) shl 16) or (((i * 2) and 0xff) shl 8) or ((i * 3) and 0xff) + } + val decoded = ThumbHashDecoder.decode(ThumbHashEncoder.encode(pixels, w, h)) + assertNotNull(decoded) + decoded!! + assertTrue( + "expected output to fit in 32x32, got ${decoded.width}x${decoded.height}", + decoded.width in 1..32 && decoded.height in 1..32, + ) + assertEquals( + "pixel buffer size must match dimensions", + decoded.width * decoded.height, + decoded.pixels.size, + ) } }