perf(thumbhash): cache cosine tables and flatten hot loop in decoder

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<Double>
  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.
This commit is contained in:
Claude
2026-04-20 00:33:26 +00:00
parent b2f297b3bb
commit 63da850e3b
3 changed files with 557 additions and 92 deletions
@@ -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)
}
}
}
@@ -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<Double>` 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<Long, DoubleArray>()
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<Double>(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)
@@ -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,
)
}
}