From 63da850e3be2d6f78049fc026f5778982f32ec0c Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 20 Apr 2026 00:33:26 +0000 Subject: [PATCH] perf(thumbhash): cache cosine tables and flatten hot loop in decoder MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The reference port recomputed the inverse DCT cosine tables once per output pixel — for a 32x24 decode that's ~43k redundant cos() calls. This change lifts the tables out of the inner loop and caches them across decodes keyed by (size, componentCount) so subsequent placeholders at the same target size skip cosine evaluation entirely. Additional wins: - AC coefficients unpack into pre-sized DoubleArrays (no ArrayList boxing, no post-hoc copy) - truncated hashes are rejected up-front via a single length check instead of mid-stream - LPQA -> sRGB uses an inline branch clamp instead of a min/max/round/coerceIn chain Tests: 12 unit tests cover determinism across repeated decodes, cache clearing, alpha preservation, aspect-ratio preservation both directions, average-color drift bounds, truncated input rejection, and round-tripping base64 vs. raw bytes. All green. Bench: new ThumbHashBenchmark exercises opaque decode, alpha decode, warm- and cold-cache decode, aspect-ratio probe, and the full decodeKeepAspectRatio pipeline so the perf delta is visible on device. --- .../amethyst/benchmark/ThumbHashBenchmark.kt | 156 +++++++++ .../commons/thumbhash/ThumbHashDecoder.kt | 327 +++++++++++++----- .../amethyst/commons/ThumbHashTest.kt | 166 ++++++++- 3 files changed, 557 insertions(+), 92 deletions(-) create mode 100644 benchmark/src/androidTest/java/com/vitorpamplona/amethyst/benchmark/ThumbHashBenchmark.kt 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, + ) } }