package com.infendro.hash.keccak import com.infendro.bytes.ByteOrder.LITTLE_ENDIAN import com.infendro.bytes.bytearray.copyInto import com.infendro.bytes.long.copyInto import com.infendro.hash.HashFunction import com.infendro.hash.util.align open class Keccak( override val bytes: Int, private val rate: Int, private val domain: Byte, ) : HashFunction { init { require(bytes > 0) require(rate in 8..192 step 8) } override val block: Int get() = rate private val rc = ulongArrayOf( 0x0000000000000001UL, 0x0000000000008082UL, 0x800000000000808aUL, 0x8000000080008000UL, 0x000000000000808bUL, 0x0000000080000001UL, 0x8000000080008081UL, 0x8000000000008009UL, 0x000000000000008aUL, 0x0000000000000088UL, 0x0000000080008009UL, 0x000000008000000aUL, 0x000000008000808bUL, 0x800000000000008bUL, 0x8000000000008089UL, 0x8000000000008003UL, 0x8000000000008002UL, 0x8000000000000080UL, 0x000000000000800aUL, 0x800000008000000aUL, 0x8000000080008081UL, 0x8000000000008080UL, 0x0000000080000001UL, 0x8000000080008008UL, ).asLongArray() private val rotation = arrayOf( intArrayOf(0, 36, 3, 41, 18), intArrayOf(1, 44, 10, 45, 2), intArrayOf(62, 6, 43, 15, 61), intArrayOf(28, 55, 25, 21, 56), intArrayOf(27, 20, 39, 8, 14), ) override fun hashInto(value: ByteArray, destination: ByteArray, destinationOffset: Int) { val input = pad(value) val state = LongArray(25) val w = LongArray(block / 8) val b = LongArray(25) val c = LongArray(5) val d = LongArray(5) for (i in input.indices step block) { absorb(state, input, i, w) permute(state, b, c, d) } squeeze(state, b, c, d, destination, destinationOffset) } private fun pad(value: ByteArray): ByteArray { val size = (value.size + 2).align(block) val result = ByteArray(size) value.copyInto(result) result[value.size] = domain result[result.lastIndex] = 0x80U.toByte() return result } private fun absorb(state: LongArray, value: ByteArray, offset: Int, w: LongArray) { value.copyInto(w, startIndex = offset, endIndex = offset + block, order = LITTLE_ENDIAN) for (i in w.indices) { state[i] = state[i] xor w[i] } } private fun squeeze( state: LongArray, b: LongArray, c: LongArray, d: LongArray, destination: ByteArray, destinationOffset: Int, ) { val buffer = ByteArray(8) val threshold = rate / 8 var consumed = 0 var i = 0 while (consumed < bytes) { if (i == threshold) { permute(state, b, c, d) i = 0 } val length = minOf(bytes - consumed, 8) state[i].copyInto(buffer, order = LITTLE_ENDIAN) buffer.copyInto(destination, destinationOffset + consumed, endIndex = length) consumed += length i++ } } private fun permute(a: LongArray, b: LongArray, c: LongArray, d: LongArray) = repeat(24) { round -> // θ for (x in 0..4) { c[x] = a[x] xor a[x + 5] xor a[x + 10] xor a[x + 15] xor a[x + 20] } for (x in 0..4) { d[x] = c[(x + 4) % 5] xor c[(x + 1) % 5].rotateLeft(1) } for (x in 0..4) for (y in 0..4) { a[x + 5 * y] = a[x + 5 * y] xor d[x] } // ρ and π for (x in 0..4) for (y in 0..4) { b[y + 5 * ((2 * x + 3 * y) % 5)] = a[x + 5 * y].rotateLeft(rotation[x][y]) } // χ for (x in 0..4) for (y in 0..4) { a[x + 5 * y] = b[x + 5 * y] xor ((b[((x + 1) % 5) + 5 * y].inv()) and b[((x + 2) % 5) + 5 * y]) } // ι a[0] = a[0] xor rc[round] } }