master xplshn/aruu / cmd / extra / b3sum.c
   1/* See LICENSE file for copyright and license details. */
   2
   3#if defined(__x86_64__) && (defined(__clang__) || (defined(__GNUC__) && !defined(__TINYC__)))
   4#define BLAKE3_X86_SIMD 1
   5#else
   6#define BLAKE3_X86_SIMD 0
   7#endif
   8
   9#include "arg.h"
  10#include "util.h"
  11
  12#include <assert.h>
  13#if BLAKE3_X86_SIMD
  14#include <immintrin.h>
  15#endif
  16#include <stdint.h>
  17#include <stdio.h>
  18#include <stdlib.h>
  19#include <string.h>
  20
  21#define BLAKE3_VERSION_STRING "1.5.0"
  22#define BLAKE3_KEY_LEN        32
  23#define BLAKE3_OUT_LEN        32
  24#define BLAKE3_BLOCK_LEN      64
  25#define BLAKE3_CHUNK_LEN      1024
  26#define BLAKE3_MAX_DEPTH      54
  27
  28#if BLAKE3_X86_SIMD
  29#define MAX_SIMD_DEGREE 16
  30#else
  31#define MAX_SIMD_DEGREE 1
  32#endif
  33
  34#define MAX_SIMD_DEGREE_OR_2 (MAX_SIMD_DEGREE > 2 ? MAX_SIMD_DEGREE : 2)
  35
  36enum Blake3Flags {
  37  CHUNK_START         = 1 << 0,
  38  CHUNK_END           = 1 << 1,
  39  PARENT              = 1 << 2,
  40  ROOT                = 1 << 3,
  41  KEYED_HASH          = 1 << 4,
  42  DERIVE_KEY_CONTEXT  = 1 << 5,
  43  DERIVE_KEY_MATERIAL = 1 << 6
  44};
  45
  46struct Blake3ChunkState {
  47  uint32_t cv[8];
  48  uint64_t chunk_counter;
  49  uint8_t  buf[BLAKE3_BLOCK_LEN];
  50  uint8_t  buf_len;
  51  uint8_t  blocks_compressed;
  52  uint8_t  flags;
  53};
  54
  55struct Blake3Hasher {
  56  uint32_t                key[8];
  57  struct Blake3ChunkState chunk;
  58  uint8_t                 cv_stack_len;
  59  uint8_t                 cv_stack[(BLAKE3_MAX_DEPTH + 1) * BLAKE3_OUT_LEN];
  60};
  61
  62void blake3_hasher_update(struct Blake3Hasher *self, const void *input, size_t input_len);
  63void blake3_hasher_finalize(const struct Blake3Hasher *self, uint8_t *out, size_t out_len);
  64
  65struct Output {
  66  uint32_t input_cv[8];
  67  uint64_t counter;
  68  uint8_t  block[BLAKE3_BLOCK_LEN];
  69  uint8_t  block_len;
  70  uint8_t  flags;
  71};
  72
  73static const uint32_t IV[8] = {
  74    0x6A09E667UL,
  75    0xBB67AE85UL,
  76    0x3C6EF372UL,
  77    0xA54FF53AUL,
  78    0x510E527FUL,
  79    0x9B05688CUL,
  80    0x1F83D9ABUL,
  81    0x5BE0CD19UL
  82};
  83
  84static const uint8_t MSG_SCHEDULE[7][16] = {
  85    {0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15},
  86    {2, 6, 3, 10, 7, 0, 4, 13, 1, 11, 12, 5, 9, 14, 15, 8},
  87    {3, 4, 10, 12, 13, 2, 7, 14, 6, 5, 9, 0, 11, 15, 8, 1},
  88    {10, 7, 12, 9, 14, 3, 13, 15, 4, 0, 11, 2, 5, 8, 1, 6},
  89    {12, 13, 9, 11, 15, 10, 14, 8, 7, 2, 5, 3, 0, 1, 6, 4},
  90    {9, 14, 11, 5, 8, 12, 15, 1, 13, 3, 0, 10, 2, 6, 4, 7},
  91    {11, 15, 5, 0, 1, 9, 8, 6, 14, 10, 2, 12, 3, 4, 7, 13},
  92};
  93
  94static inline uint32_t
  95load32(const void *src)
  96{
  97  const uint8_t *p = (const uint8_t *)src;
  98
  99  return ((uint32_t)(p[0]) << 0) | ((uint32_t)(p[1]) << 8) | ((uint32_t)(p[2]) << 16)
 100         | ((uint32_t)(p[3]) << 24);
 101}
 102
 103static inline void
 104store32(void *dst, uint32_t w)
 105{
 106  uint8_t *p = (uint8_t *)dst;
 107
 108  p[0] = (uint8_t)(w >> 0);
 109  p[1] = (uint8_t)(w >> 8);
 110  p[2] = (uint8_t)(w >> 16);
 111  p[3] = (uint8_t)(w >> 24);
 112}
 113
 114static inline void
 115store_cv_words(uint8_t bytes_out[32], uint32_t cv_words[8])
 116{
 117  store32(&bytes_out[0 * 4], cv_words[0]);
 118  store32(&bytes_out[1 * 4], cv_words[1]);
 119  store32(&bytes_out[2 * 4], cv_words[2]);
 120  store32(&bytes_out[3 * 4], cv_words[3]);
 121  store32(&bytes_out[4 * 4], cv_words[4]);
 122  store32(&bytes_out[5 * 4], cv_words[5]);
 123  store32(&bytes_out[6 * 4], cv_words[6]);
 124  store32(&bytes_out[7 * 4], cv_words[7]);
 125}
 126
 127static inline uint32_t
 128counter_low(uint64_t counter)
 129{
 130  return (uint32_t)counter;
 131}
 132
 133static inline uint32_t
 134counter_high(uint64_t counter)
 135{
 136  return (uint32_t)(counter >> 32);
 137}
 138
 139/* forward declarations */
 140#if BLAKE3_X86_SIMD
 141void blake3_compress_in_place_sse2(
 142    uint32_t      cv[8],
 143    const uint8_t block[BLAKE3_BLOCK_LEN],
 144    uint8_t       block_len,
 145    uint64_t      counter,
 146    uint8_t       flags
 147);
 148void blake3_compress_xof_sse2(
 149    const uint32_t cv[8],
 150    const uint8_t  block[BLAKE3_BLOCK_LEN],
 151    uint8_t        block_len,
 152    uint64_t       counter,
 153    uint8_t        flags,
 154    uint8_t        out[64]
 155);
 156void blake3_hash_many_sse2(
 157    const uint8_t *const *inputs,
 158    size_t                num_inputs,
 159    size_t                blocks,
 160    const uint32_t        key[8],
 161    uint64_t              counter,
 162    int                   increment_counter,
 163    uint8_t               flags,
 164    uint8_t               flags_start,
 165    uint8_t               flags_end,
 166    uint8_t              *out
 167);
 168
 169void blake3_compress_in_place_sse41(
 170    uint32_t      cv[8],
 171    const uint8_t block[BLAKE3_BLOCK_LEN],
 172    uint8_t       block_len,
 173    uint64_t      counter,
 174    uint8_t       flags
 175);
 176void blake3_compress_xof_sse41(
 177    const uint32_t cv[8],
 178    const uint8_t  block[BLAKE3_BLOCK_LEN],
 179    uint8_t        block_len,
 180    uint64_t       counter,
 181    uint8_t        flags,
 182    uint8_t        out[64]
 183);
 184void blake3_hash_many_sse41(
 185    const uint8_t *const *inputs,
 186    size_t                num_inputs,
 187    size_t                blocks,
 188    const uint32_t        key[8],
 189    uint64_t              counter,
 190    int                   increment_counter,
 191    uint8_t               flags,
 192    uint8_t               flags_start,
 193    uint8_t               flags_end,
 194    uint8_t              *out
 195);
 196
 197void blake3_hash_many_avx2(
 198    const uint8_t *const *inputs,
 199    size_t                num_inputs,
 200    size_t                blocks,
 201    const uint32_t        key[8],
 202    uint64_t              counter,
 203    int                   increment_counter,
 204    uint8_t               flags,
 205    uint8_t               flags_start,
 206    uint8_t               flags_end,
 207    uint8_t              *out
 208);
 209
 210void blake3_compress_in_place_avx512(
 211    uint32_t      cv[8],
 212    const uint8_t block[BLAKE3_BLOCK_LEN],
 213    uint8_t       block_len,
 214    uint64_t      counter,
 215    uint8_t       flags
 216);
 217void blake3_compress_xof_avx512(
 218    const uint32_t cv[8],
 219    const uint8_t  block[BLAKE3_BLOCK_LEN],
 220    uint8_t        block_len,
 221    uint64_t       counter,
 222    uint8_t        flags,
 223    uint8_t        out[64]
 224);
 225void blake3_hash_many_avx512(
 226    const uint8_t *const *inputs,
 227    size_t                num_inputs,
 228    size_t                blocks,
 229    const uint32_t        key[8],
 230    uint64_t              counter,
 231    int                   increment_counter,
 232    uint8_t               flags,
 233    uint8_t               flags_start,
 234    uint8_t               flags_end,
 235    uint8_t              *out
 236);
 237#endif
 238
 239void blake3_compress_in_place_portable(
 240    uint32_t      cv[8],
 241    const uint8_t block[BLAKE3_BLOCK_LEN],
 242    uint8_t       block_len,
 243    uint64_t      counter,
 244    uint8_t       flags
 245);
 246void blake3_compress_xof_portable(
 247    const uint32_t cv[8],
 248    const uint8_t  block[BLAKE3_BLOCK_LEN],
 249    uint8_t        block_len,
 250    uint64_t       counter,
 251    uint8_t        flags,
 252    uint8_t        out[64]
 253);
 254void blake3_hash_many_portable(
 255    const uint8_t *const *inputs,
 256    size_t                num_inputs,
 257    size_t                blocks,
 258    const uint32_t        key[8],
 259    uint64_t              counter,
 260    int                   increment_counter,
 261    uint8_t               flags,
 262    uint8_t               flags_start,
 263    uint8_t               flags_end,
 264    uint8_t              *out
 265);
 266
 267void blake3_compress_in_place(
 268    uint32_t      cv[8],
 269    const uint8_t block[BLAKE3_BLOCK_LEN],
 270    uint8_t       block_len,
 271    uint64_t      counter,
 272    uint8_t       flags
 273);
 274void blake3_compress_xof(
 275    const uint32_t cv[8],
 276    const uint8_t  block[BLAKE3_BLOCK_LEN],
 277    uint8_t        block_len,
 278    uint64_t       counter,
 279    uint8_t        flags,
 280    uint8_t        out[64]
 281);
 282void blake3_hash_many(
 283    const uint8_t *const *inputs,
 284    size_t                num_inputs,
 285    size_t                blocks,
 286    const uint32_t        key[8],
 287    uint64_t              counter,
 288    int                   increment_counter,
 289    uint8_t               flags,
 290    uint8_t               flags_start,
 291    uint8_t               flags_end,
 292    uint8_t              *out
 293);
 294
 295/* portable implementations */
 296static inline uint32_t
 297rotr32(uint32_t w, uint32_t c)
 298{
 299  return (w >> c) | (w << (32 - c));
 300}
 301
 302static inline void
 303g_portable(uint32_t *state, size_t a, size_t b, size_t c, size_t d, uint32_t x, uint32_t y)
 304{
 305  state[a] = state[a] + state[b] + x;
 306  state[d] = rotr32(state[d] ^ state[a], 16);
 307  state[c] = state[c] + state[d];
 308  state[b] = rotr32(state[b] ^ state[c], 12);
 309  state[a] = state[a] + state[b] + y;
 310  state[d] = rotr32(state[d] ^ state[a], 8);
 311  state[c] = state[c] + state[d];
 312  state[b] = rotr32(state[b] ^ state[c], 7);
 313}
 314
 315static inline void
 316round_fn_portable(uint32_t state[16], const uint32_t *msg, size_t round)
 317{
 318  const uint8_t *schedule = MSG_SCHEDULE[round];
 319
 320  g_portable(state, 0, 4, 8, 12, msg[schedule[0]], msg[schedule[1]]);
 321  g_portable(state, 1, 5, 9, 13, msg[schedule[2]], msg[schedule[3]]);
 322  g_portable(state, 2, 6, 10, 14, msg[schedule[4]], msg[schedule[5]]);
 323  g_portable(state, 3, 7, 11, 15, msg[schedule[6]], msg[schedule[7]]);
 324
 325  g_portable(state, 0, 5, 10, 15, msg[schedule[8]], msg[schedule[9]]);
 326  g_portable(state, 1, 6, 11, 12, msg[schedule[10]], msg[schedule[11]]);
 327  g_portable(state, 2, 7, 8, 13, msg[schedule[12]], msg[schedule[13]]);
 328  g_portable(state, 3, 4, 9, 14, msg[schedule[14]], msg[schedule[15]]);
 329}
 330
 331static inline void
 332compress_pre_portable(
 333    uint32_t       state[16],
 334    const uint32_t cv[8],
 335    const uint8_t  block[BLAKE3_BLOCK_LEN],
 336    uint8_t        block_len,
 337    uint64_t       counter,
 338    uint8_t        flags
 339)
 340{
 341  uint32_t block_words[16];
 342
 343  block_words[0]  = load32(block + 4 * 0);
 344  block_words[1]  = load32(block + 4 * 1);
 345  block_words[2]  = load32(block + 4 * 2);
 346  block_words[3]  = load32(block + 4 * 3);
 347  block_words[4]  = load32(block + 4 * 4);
 348  block_words[5]  = load32(block + 4 * 5);
 349  block_words[6]  = load32(block + 4 * 6);
 350  block_words[7]  = load32(block + 4 * 7);
 351  block_words[8]  = load32(block + 4 * 8);
 352  block_words[9]  = load32(block + 4 * 9);
 353  block_words[10] = load32(block + 4 * 10);
 354  block_words[11] = load32(block + 4 * 11);
 355  block_words[12] = load32(block + 4 * 12);
 356  block_words[13] = load32(block + 4 * 13);
 357  block_words[14] = load32(block + 4 * 14);
 358  block_words[15] = load32(block + 4 * 15);
 359
 360  state[0]  = cv[0];
 361  state[1]  = cv[1];
 362  state[2]  = cv[2];
 363  state[3]  = cv[3];
 364  state[4]  = cv[4];
 365  state[5]  = cv[5];
 366  state[6]  = cv[6];
 367  state[7]  = cv[7];
 368  state[8]  = IV[0];
 369  state[9]  = IV[1];
 370  state[10] = IV[2];
 371  state[11] = IV[3];
 372  state[12] = counter_low(counter);
 373  state[13] = counter_high(counter);
 374  state[14] = (uint32_t)block_len;
 375  state[15] = (uint32_t)flags;
 376
 377  round_fn_portable(state, &block_words[0], 0);
 378  round_fn_portable(state, &block_words[0], 1);
 379  round_fn_portable(state, &block_words[0], 2);
 380  round_fn_portable(state, &block_words[0], 3);
 381  round_fn_portable(state, &block_words[0], 4);
 382  round_fn_portable(state, &block_words[0], 5);
 383  round_fn_portable(state, &block_words[0], 6);
 384}
 385
 386void
 387blake3_compress_in_place_portable(
 388    uint32_t      cv[8],
 389    const uint8_t block[BLAKE3_BLOCK_LEN],
 390    uint8_t       block_len,
 391    uint64_t      counter,
 392    uint8_t       flags
 393)
 394{
 395  uint32_t state[16];
 396
 397  compress_pre_portable(state, cv, block, block_len, counter, flags);
 398  cv[0] = state[0] ^ state[8];
 399  cv[1] = state[1] ^ state[9];
 400  cv[2] = state[2] ^ state[10];
 401  cv[3] = state[3] ^ state[11];
 402  cv[4] = state[4] ^ state[12];
 403  cv[5] = state[5] ^ state[13];
 404  cv[6] = state[6] ^ state[14];
 405  cv[7] = state[7] ^ state[15];
 406}
 407
 408void
 409blake3_compress_xof_portable(
 410    const uint32_t cv[8],
 411    const uint8_t  block[BLAKE3_BLOCK_LEN],
 412    uint8_t        block_len,
 413    uint64_t       counter,
 414    uint8_t        flags,
 415    uint8_t        out[64]
 416)
 417{
 418  uint32_t state[16];
 419
 420  compress_pre_portable(state, cv, block, block_len, counter, flags);
 421
 422  store32(&out[0 * 4], state[0] ^ state[8]);
 423  store32(&out[1 * 4], state[1] ^ state[9]);
 424  store32(&out[2 * 4], state[2] ^ state[10]);
 425  store32(&out[3 * 4], state[3] ^ state[11]);
 426  store32(&out[4 * 4], state[4] ^ state[12]);
 427  store32(&out[5 * 4], state[5] ^ state[13]);
 428  store32(&out[6 * 4], state[6] ^ state[14]);
 429  store32(&out[7 * 4], state[7] ^ state[15]);
 430  store32(&out[8 * 4], state[8] ^ cv[0]);
 431  store32(&out[9 * 4], state[9] ^ cv[1]);
 432  store32(&out[10 * 4], state[10] ^ cv[2]);
 433  store32(&out[11 * 4], state[11] ^ cv[3]);
 434  store32(&out[12 * 4], state[12] ^ cv[4]);
 435  store32(&out[13 * 4], state[13] ^ cv[5]);
 436  store32(&out[14 * 4], state[14] ^ cv[6]);
 437  store32(&out[15 * 4], state[15] ^ cv[7]);
 438}
 439
 440static inline void
 441hash_one_portable(
 442    const uint8_t *input,
 443    size_t         blocks,
 444    const uint32_t key[8],
 445    uint64_t       counter,
 446    uint8_t        flags,
 447    uint8_t        flags_start,
 448    uint8_t        flags_end,
 449    uint8_t        out[BLAKE3_OUT_LEN]
 450)
 451{
 452  uint32_t cv[8];
 453  uint8_t  block_flags;
 454
 455  memcpy(cv, key, BLAKE3_KEY_LEN);
 456  block_flags = flags | flags_start;
 457  while (blocks > 0) {
 458    if (blocks == 1) {
 459      block_flags |= flags_end;
 460    }
 461    blake3_compress_in_place_portable(cv, input, BLAKE3_BLOCK_LEN, counter, block_flags);
 462    input = &input[BLAKE3_BLOCK_LEN];
 463    blocks -= 1;
 464    block_flags = flags;
 465  }
 466  store_cv_words(out, cv);
 467}
 468
 469void
 470blake3_hash_many_portable(
 471    const uint8_t *const *inputs,
 472    size_t                num_inputs,
 473    size_t                blocks,
 474    const uint32_t        key[8],
 475    uint64_t              counter,
 476    int                   increment_counter,
 477    uint8_t               flags,
 478    uint8_t               flags_start,
 479    uint8_t               flags_end,
 480    uint8_t              *out
 481)
 482{
 483  while (num_inputs > 0) {
 484    hash_one_portable(inputs[0], blocks, key, counter, flags, flags_start, flags_end, out);
 485    if (increment_counter) {
 486      counter += 1;
 487    }
 488    inputs += 1;
 489    num_inputs -= 1;
 490    out = &out[BLAKE3_OUT_LEN];
 491  }
 492}
 493
 494/* cpu features detection */
 495enum { SSE2 = 1 << 0, SSE41 = 1 << 1, AVX2 = 1 << 2, AVX512 = 1 << 3 };
 496
 497static int blake3_cpu_features = 0;
 498static int blake3_cpu_detected = 0;
 499
 500#if BLAKE3_X86_SIMD
 501#include <cpuid.h>
 502
 503static void
 504blake3_cpuid(uint32_t out[4], uint32_t id, uint32_t sid)
 505{
 506  __cpuid_count(id, sid, out[0], out[1], out[2], out[3]);
 507}
 508
 509static uint64_t
 510blake3_xgetbv(void)
 511{
 512  uint32_t eax, edx;
 513
 514  __asm__ volatile("xgetbv" : "=a"(eax), "=d"(edx) : "c"(0));
 515  return ((uint64_t)edx << 32) | eax;
 516}
 517#endif
 518
 519static void
 520blake3_detect_cpu_features(void)
 521{
 522#if BLAKE3_X86_SIMD
 523  enum { EAX, EBX, ECX, EDX };
 524  uint32_t regs[4];
 525  uint64_t xcr0;
 526  int      features = 0;
 527
 528  blake3_cpuid(regs, 1, 0);
 529  if (regs[EDX] & (1UL << 26))
 530    features |= SSE2;
 531  if (regs[ECX] & (1UL << 19))
 532    features |= SSE41;
 533  /* osxsave */
 534  if (regs[ECX] & (1UL << 27)) {
 535    blake3_cpuid(regs, 0, 0);
 536    if (regs[EAX] >= 7) {
 537      blake3_cpuid(regs, 7, 0);
 538      xcr0 = blake3_xgetbv();
 539      /* avx2 and xcr0 sse, avx */
 540      if ((regs[EBX] & (1UL << 5)) && (xcr0 & 0x06) == 0x06)
 541        features |= AVX2;
 542      /* avx512f, avx512vl and xcr0 opmask, zmm_hi256,
 543       * hi16_zmm */
 544      if ((regs[EBX] & (1UL << 31 | 1UL << 16)) && (xcr0 & 0xe0) == 0xe0)
 545        features |= AVX512;
 546    }
 547  }
 548  blake3_cpu_features = features;
 549#endif
 550  blake3_cpu_detected = 1;
 551}
 552
 553#if BLAKE3_X86_SIMD
 554__attribute__((constructor)) static void
 555blake3_init_cpu(void)
 556{
 557  if (!blake3_cpu_detected)
 558    blake3_detect_cpu_features();
 559}
 560#endif
 561
 562#if BLAKE3_X86_SIMD
 563#if defined(__clang__)
 564#pragma clang attribute push(__attribute__((target("sse2"))), apply_to = function)
 565#elif defined(__GNUC__)
 566#pragma GCC push_options
 567#pragma GCC target("sse2")
 568#endif
 569#define DEGREE_SSE2 4
 570
 571#define _mm_shuffle_ps2(a, b, c)                                                                   \
 572  (_mm_castps_si128(_mm_shuffle_ps(_mm_castsi128_ps(a), _mm_castsi128_ps(b), (c))))
 573
 574static inline __m128i
 575loadu_sse2(const uint8_t src[16])
 576{
 577  return _mm_loadu_si128((const __m128i *)src);
 578}
 579
 580static inline void
 581storeu_sse2(__m128i src, uint8_t dest[16])
 582{
 583  _mm_storeu_si128((__m128i *)dest, src);
 584}
 585
 586static inline __m128i
 587addv_sse2(__m128i a, __m128i b)
 588{
 589  return _mm_add_epi32(a, b);
 590}
 591
 592static inline __m128i
 593xorv_sse2(__m128i a, __m128i b)
 594{
 595  return _mm_xor_si128(a, b);
 596}
 597
 598static inline __m128i
 599set1_sse2(uint32_t x)
 600{
 601  return _mm_set1_epi32((int32_t)x);
 602}
 603
 604static inline __m128i
 605set4_sse2(uint32_t a, uint32_t b, uint32_t c, uint32_t d)
 606{
 607  return _mm_setr_epi32((int32_t)a, (int32_t)b, (int32_t)c, (int32_t)d);
 608}
 609
 610static inline __m128i
 611rot16_sse2(__m128i x)
 612{
 613  return _mm_shufflehi_epi16(_mm_shufflelo_epi16(x, 0xB1), 0xB1);
 614}
 615
 616static inline __m128i
 617rot12_sse2(__m128i x)
 618{
 619  return xorv_sse2(_mm_srli_epi32(x, 12), _mm_slli_epi32(x, 32 - 12));
 620}
 621
 622static inline __m128i
 623rot8_sse2(__m128i x)
 624{
 625  return xorv_sse2(_mm_srli_epi32(x, 8), _mm_slli_epi32(x, 32 - 8));
 626}
 627
 628static inline __m128i
 629rot7_sse2(__m128i x)
 630{
 631  return xorv_sse2(_mm_srli_epi32(x, 7), _mm_slli_epi32(x, 32 - 7));
 632}
 633
 634static inline void
 635g1_sse2(__m128i *row0, __m128i *row1, __m128i *row2, __m128i *row3, __m128i m)
 636{
 637  *row0 = addv_sse2(addv_sse2(*row0, m), *row1);
 638  *row3 = xorv_sse2(*row3, *row0);
 639  *row3 = rot16_sse2(*row3);
 640  *row2 = addv_sse2(*row2, *row3);
 641  *row1 = xorv_sse2(*row1, *row2);
 642  *row1 = rot12_sse2(*row1);
 643}
 644
 645static inline void
 646g2_sse2(__m128i *row0, __m128i *row1, __m128i *row2, __m128i *row3, __m128i m)
 647{
 648  *row0 = addv_sse2(addv_sse2(*row0, m), *row1);
 649  *row3 = xorv_sse2(*row3, *row0);
 650  *row3 = rot8_sse2(*row3);
 651  *row2 = addv_sse2(*row2, *row3);
 652  *row1 = xorv_sse2(*row1, *row2);
 653  *row1 = rot7_sse2(*row1);
 654}
 655
 656static inline void
 657diagonalize_sse2(__m128i *row0, __m128i *row2, __m128i *row3)
 658{
 659  *row0 = _mm_shuffle_epi32(*row0, _MM_SHUFFLE(2, 1, 0, 3));
 660  *row3 = _mm_shuffle_epi32(*row3, _MM_SHUFFLE(1, 0, 3, 2));
 661  *row2 = _mm_shuffle_epi32(*row2, _MM_SHUFFLE(0, 3, 2, 1));
 662}
 663
 664static inline void
 665undiagonalize_sse2(__m128i *row0, __m128i *row2, __m128i *row3)
 666{
 667  *row0 = _mm_shuffle_epi32(*row0, _MM_SHUFFLE(0, 3, 2, 1));
 668  *row3 = _mm_shuffle_epi32(*row3, _MM_SHUFFLE(1, 0, 3, 2));
 669  *row2 = _mm_shuffle_epi32(*row2, _MM_SHUFFLE(2, 1, 0, 3));
 670}
 671
 672static inline __m128i
 673blend_epi16_sse2(__m128i a, __m128i b, const int16_t imm8)
 674{
 675  const __m128i bits = _mm_set_epi16(0x80, 0x40, 0x20, 0x10, 0x08, 0x04, 0x02, 0x01);
 676  __m128i       mask = _mm_set1_epi16(imm8);
 677
 678  mask = _mm_and_si128(mask, bits);
 679  mask = _mm_cmpeq_epi16(mask, bits);
 680  return _mm_or_si128(_mm_and_si128(mask, b), _mm_andnot_si128(mask, a));
 681}
 682
 683static inline void
 684compress_pre_sse2(
 685    __m128i        rows[4],
 686    const uint32_t cv[8],
 687    const uint8_t  block[BLAKE3_BLOCK_LEN],
 688    uint8_t        block_len,
 689    uint64_t       counter,
 690    uint8_t        flags
 691)
 692{
 693  __m128i m0, m1, m2, m3;
 694  __m128i t0, t1, t2, t3, tt;
 695
 696  rows[0] = loadu_sse2((uint8_t *)&cv[0]);
 697  rows[1] = loadu_sse2((uint8_t *)&cv[4]);
 698  rows[2] = set4_sse2(IV[0], IV[1], IV[2], IV[3]);
 699  rows[3] =
 700      set4_sse2(counter_low(counter), counter_high(counter), (uint32_t)block_len, (uint32_t)flags);
 701
 702  m0 = loadu_sse2(&block[sizeof(__m128i) * 0]);
 703  m1 = loadu_sse2(&block[sizeof(__m128i) * 1]);
 704  m2 = loadu_sse2(&block[sizeof(__m128i) * 2]);
 705  m3 = loadu_sse2(&block[sizeof(__m128i) * 3]);
 706
 707  /* round 1 */
 708  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(2, 0, 2, 0));
 709  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 710  t1 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 3, 1));
 711  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 712  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 713  t2 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(2, 0, 2, 0));
 714  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(2, 1, 0, 3));
 715  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 716  t3 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 1, 3, 1));
 717  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(2, 1, 0, 3));
 718  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 719  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 720  m0 = t0;
 721  m1 = t1;
 722  m2 = t2;
 723  m3 = t3;
 724
 725  /* round 2 */
 726  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
 727  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
 728  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 729  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
 730  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
 731  t1 = blend_epi16_sse2(tt, t1, 0xCC);
 732  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 733  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 734  t2 = _mm_unpacklo_epi64(m3, m1);
 735  tt = blend_epi16_sse2(t2, m2, 0xC0);
 736  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
 737  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 738  t3 = _mm_unpackhi_epi32(m1, m3);
 739  tt = _mm_unpacklo_epi32(m2, t3);
 740  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
 741  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 742  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 743  m0 = t0;
 744  m1 = t1;
 745  m2 = t2;
 746  m3 = t3;
 747
 748  /* round 3 */
 749  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
 750  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
 751  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 752  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
 753  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
 754  t1 = blend_epi16_sse2(tt, t1, 0xCC);
 755  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 756  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 757  t2 = _mm_unpacklo_epi64(m3, m1);
 758  tt = blend_epi16_sse2(t2, m2, 0xC0);
 759  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
 760  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 761  t3 = _mm_unpackhi_epi32(m1, m3);
 762  tt = _mm_unpacklo_epi32(m2, t3);
 763  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
 764  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 765  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 766  m0 = t0;
 767  m1 = t1;
 768  m2 = t2;
 769  m3 = t3;
 770
 771  /* round 4 */
 772  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
 773  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
 774  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 775  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
 776  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
 777  t1 = blend_epi16_sse2(tt, t1, 0xCC);
 778  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 779  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 780  t2 = _mm_unpacklo_epi64(m3, m1);
 781  tt = blend_epi16_sse2(t2, m2, 0xC0);
 782  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
 783  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 784  t3 = _mm_unpackhi_epi32(m1, m3);
 785  tt = _mm_unpacklo_epi32(m2, t3);
 786  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
 787  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 788  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 789  m0 = t0;
 790  m1 = t1;
 791  m2 = t2;
 792  m3 = t3;
 793
 794  /* round 5 */
 795  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
 796  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
 797  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 798  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
 799  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
 800  t1 = blend_epi16_sse2(tt, t1, 0xCC);
 801  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 802  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 803  t2 = _mm_unpacklo_epi64(m3, m1);
 804  tt = blend_epi16_sse2(t2, m2, 0xC0);
 805  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
 806  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 807  t3 = _mm_unpackhi_epi32(m1, m3);
 808  tt = _mm_unpacklo_epi32(m2, t3);
 809  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
 810  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 811  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 812  m0 = t0;
 813  m1 = t1;
 814  m2 = t2;
 815  m3 = t3;
 816
 817  /* round 6 */
 818  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
 819  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
 820  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 821  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
 822  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
 823  t1 = blend_epi16_sse2(tt, t1, 0xCC);
 824  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 825  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 826  t2 = _mm_unpacklo_epi64(m3, m1);
 827  tt = blend_epi16_sse2(t2, m2, 0xC0);
 828  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
 829  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 830  t3 = _mm_unpackhi_epi32(m1, m3);
 831  tt = _mm_unpacklo_epi32(m2, t3);
 832  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
 833  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 834  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 835  m0 = t0;
 836  m1 = t1;
 837  m2 = t2;
 838  m3 = t3;
 839
 840  /* round 7 */
 841  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
 842  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
 843  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t0);
 844  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
 845  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
 846  t1 = blend_epi16_sse2(tt, t1, 0xCC);
 847  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t1);
 848  diagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 849  t2 = _mm_unpacklo_epi64(m3, m1);
 850  tt = blend_epi16_sse2(t2, m2, 0xC0);
 851  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
 852  g1_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t2);
 853  t3 = _mm_unpackhi_epi32(m1, m3);
 854  tt = _mm_unpacklo_epi32(m2, t3);
 855  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
 856  g2_sse2(&rows[0], &rows[1], &rows[2], &rows[3], t3);
 857  undiagonalize_sse2(&rows[0], &rows[2], &rows[3]);
 858}
 859
 860void
 861blake3_compress_in_place_sse2(
 862    uint32_t      cv[8],
 863    const uint8_t block[BLAKE3_BLOCK_LEN],
 864    uint8_t       block_len,
 865    uint64_t      counter,
 866    uint8_t       flags
 867)
 868{
 869  __m128i rows[4];
 870
 871  compress_pre_sse2(rows, cv, block, block_len, counter, flags);
 872  storeu_sse2(xorv_sse2(rows[0], rows[2]), (uint8_t *)&cv[0]);
 873  storeu_sse2(xorv_sse2(rows[1], rows[3]), (uint8_t *)&cv[4]);
 874}
 875
 876void
 877blake3_compress_xof_sse2(
 878    const uint32_t cv[8],
 879    const uint8_t  block[BLAKE3_BLOCK_LEN],
 880    uint8_t        block_len,
 881    uint64_t       counter,
 882    uint8_t        flags,
 883    uint8_t        out_buf[64]
 884)
 885{
 886  __m128i rows[4];
 887
 888  compress_pre_sse2(rows, cv, block, block_len, counter, flags);
 889  storeu_sse2(xorv_sse2(rows[0], rows[2]), &out_buf[0]);
 890  storeu_sse2(xorv_sse2(rows[1], rows[3]), &out_buf[16]);
 891  storeu_sse2(xorv_sse2(rows[2], loadu_sse2((uint8_t *)&cv[0])), &out_buf[32]);
 892  storeu_sse2(xorv_sse2(rows[3], loadu_sse2((uint8_t *)&cv[4])), &out_buf[48]);
 893}
 894
 895static inline void
 896round_fn_sse2(__m128i v[16], __m128i m[16], size_t r)
 897{
 898  v[0]  = addv_sse2(v[0], m[(size_t)MSG_SCHEDULE[r][0]]);
 899  v[1]  = addv_sse2(v[1], m[(size_t)MSG_SCHEDULE[r][2]]);
 900  v[2]  = addv_sse2(v[2], m[(size_t)MSG_SCHEDULE[r][4]]);
 901  v[3]  = addv_sse2(v[3], m[(size_t)MSG_SCHEDULE[r][6]]);
 902  v[0]  = addv_sse2(v[0], v[4]);
 903  v[1]  = addv_sse2(v[1], v[5]);
 904  v[2]  = addv_sse2(v[2], v[6]);
 905  v[3]  = addv_sse2(v[3], v[7]);
 906  v[12] = xorv_sse2(v[12], v[0]);
 907  v[13] = xorv_sse2(v[13], v[1]);
 908  v[14] = xorv_sse2(v[14], v[2]);
 909  v[15] = xorv_sse2(v[15], v[3]);
 910  v[12] = rot16_sse2(v[12]);
 911  v[13] = rot16_sse2(v[13]);
 912  v[14] = rot16_sse2(v[14]);
 913  v[15] = rot16_sse2(v[15]);
 914  v[8]  = addv_sse2(v[8], v[12]);
 915  v[9]  = addv_sse2(v[9], v[13]);
 916  v[10] = addv_sse2(v[10], v[14]);
 917  v[11] = addv_sse2(v[11], v[15]);
 918  v[4]  = xorv_sse2(v[4], v[8]);
 919  v[5]  = xorv_sse2(v[5], v[9]);
 920  v[6]  = xorv_sse2(v[6], v[10]);
 921  v[7]  = xorv_sse2(v[7], v[11]);
 922  v[4]  = rot12_sse2(v[4]);
 923  v[5]  = rot12_sse2(v[5]);
 924  v[6]  = rot12_sse2(v[6]);
 925  v[7]  = rot12_sse2(v[7]);
 926  v[0]  = addv_sse2(v[0], m[(size_t)MSG_SCHEDULE[r][1]]);
 927  v[1]  = addv_sse2(v[1], m[(size_t)MSG_SCHEDULE[r][3]]);
 928  v[2]  = addv_sse2(v[2], m[(size_t)MSG_SCHEDULE[r][5]]);
 929  v[3]  = addv_sse2(v[3], m[(size_t)MSG_SCHEDULE[r][7]]);
 930  v[0]  = addv_sse2(v[0], v[4]);
 931  v[1]  = addv_sse2(v[1], v[5]);
 932  v[2]  = addv_sse2(v[2], v[6]);
 933  v[3]  = addv_sse2(v[3], v[7]);
 934  v[12] = xorv_sse2(v[12], v[0]);
 935  v[13] = xorv_sse2(v[13], v[1]);
 936  v[14] = xorv_sse2(v[14], v[2]);
 937  v[15] = xorv_sse2(v[15], v[3]);
 938  v[12] = rot8_sse2(v[12]);
 939  v[13] = rot8_sse2(v[13]);
 940  v[14] = rot8_sse2(v[14]);
 941  v[15] = rot8_sse2(v[15]);
 942  v[8]  = addv_sse2(v[8], v[12]);
 943  v[9]  = addv_sse2(v[9], v[13]);
 944  v[10] = addv_sse2(v[10], v[14]);
 945  v[11] = addv_sse2(v[11], v[15]);
 946  v[4]  = xorv_sse2(v[4], v[8]);
 947  v[5]  = xorv_sse2(v[5], v[9]);
 948  v[6]  = xorv_sse2(v[6], v[10]);
 949  v[7]  = xorv_sse2(v[7], v[11]);
 950  v[4]  = rot7_sse2(v[4]);
 951  v[5]  = rot7_sse2(v[5]);
 952  v[6]  = rot7_sse2(v[6]);
 953  v[7]  = rot7_sse2(v[7]);
 954
 955  v[0]  = addv_sse2(v[0], m[(size_t)MSG_SCHEDULE[r][8]]);
 956  v[1]  = addv_sse2(v[1], m[(size_t)MSG_SCHEDULE[r][10]]);
 957  v[2]  = addv_sse2(v[2], m[(size_t)MSG_SCHEDULE[r][12]]);
 958  v[3]  = addv_sse2(v[3], m[(size_t)MSG_SCHEDULE[r][14]]);
 959  v[0]  = addv_sse2(v[0], v[5]);
 960  v[1]  = addv_sse2(v[1], v[6]);
 961  v[2]  = addv_sse2(v[2], v[7]);
 962  v[3]  = addv_sse2(v[3], v[4]);
 963  v[15] = xorv_sse2(v[15], v[0]);
 964  v[12] = xorv_sse2(v[12], v[1]);
 965  v[13] = xorv_sse2(v[13], v[2]);
 966  v[14] = xorv_sse2(v[14], v[3]);
 967  v[15] = rot16_sse2(v[15]);
 968  v[12] = rot16_sse2(v[12]);
 969  v[13] = rot16_sse2(v[13]);
 970  v[14] = rot16_sse2(v[14]);
 971  v[10] = addv_sse2(v[10], v[15]);
 972  v[11] = addv_sse2(v[11], v[12]);
 973  v[8]  = addv_sse2(v[8], v[13]);
 974  v[9]  = addv_sse2(v[9], v[14]);
 975  v[5]  = xorv_sse2(v[5], v[10]);
 976  v[6]  = xorv_sse2(v[6], v[11]);
 977  v[7]  = xorv_sse2(v[7], v[8]);
 978  v[4]  = xorv_sse2(v[4], v[9]);
 979  v[5]  = rot12_sse2(v[5]);
 980  v[6]  = rot12_sse2(v[6]);
 981  v[7]  = rot12_sse2(v[7]);
 982  v[4]  = rot12_sse2(v[4]);
 983  v[0]  = addv_sse2(v[0], m[(size_t)MSG_SCHEDULE[r][9]]);
 984  v[1]  = addv_sse2(v[1], m[(size_t)MSG_SCHEDULE[r][11]]);
 985  v[2]  = addv_sse2(v[2], m[(size_t)MSG_SCHEDULE[r][13]]);
 986  v[3]  = addv_sse2(v[3], m[(size_t)MSG_SCHEDULE[r][15]]);
 987  v[0]  = addv_sse2(v[0], v[5]);
 988  v[1]  = addv_sse2(v[1], v[6]);
 989  v[2]  = addv_sse2(v[2], v[7]);
 990  v[3]  = addv_sse2(v[3], v[4]);
 991  v[15] = xorv_sse2(v[15], v[0]);
 992  v[12] = xorv_sse2(v[12], v[1]);
 993  v[13] = xorv_sse2(v[13], v[2]);
 994  v[14] = xorv_sse2(v[14], v[3]);
 995  v[15] = rot8_sse2(v[15]);
 996  v[12] = rot8_sse2(v[12]);
 997  v[13] = rot8_sse2(v[13]);
 998  v[14] = rot8_sse2(v[14]);
 999  v[10] = addv_sse2(v[10], v[15]);
1000  v[11] = addv_sse2(v[11], v[12]);
1001  v[8]  = addv_sse2(v[8], v[13]);
1002  v[9]  = addv_sse2(v[9], v[14]);
1003  v[5]  = xorv_sse2(v[5], v[10]);
1004  v[6]  = xorv_sse2(v[6], v[11]);
1005  v[7]  = xorv_sse2(v[7], v[8]);
1006  v[4]  = xorv_sse2(v[4], v[9]);
1007  v[5]  = rot7_sse2(v[5]);
1008  v[6]  = rot7_sse2(v[6]);
1009  v[7]  = rot7_sse2(v[7]);
1010  v[4]  = rot7_sse2(v[4]);
1011}
1012
1013static inline void
1014transpose_vecs_sse2(__m128i vecs[DEGREE_SSE2])
1015{
1016  __m128i ab_01 = _mm_unpacklo_epi32(vecs[0], vecs[1]);
1017  __m128i ab_23 = _mm_unpackhi_epi32(vecs[0], vecs[1]);
1018  __m128i cd_01 = _mm_unpacklo_epi32(vecs[2], vecs[3]);
1019  __m128i cd_23 = _mm_unpackhi_epi32(vecs[2], vecs[3]);
1020
1021  __m128i abcd_0 = _mm_unpacklo_epi64(ab_01, cd_01);
1022  __m128i abcd_1 = _mm_unpackhi_epi64(ab_01, cd_01);
1023  __m128i abcd_2 = _mm_unpacklo_epi64(ab_23, cd_23);
1024  __m128i abcd_3 = _mm_unpackhi_epi64(ab_23, cd_23);
1025
1026  vecs[0] = abcd_0;
1027  vecs[1] = abcd_1;
1028  vecs[2] = abcd_2;
1029  vecs[3] = abcd_3;
1030}
1031
1032static inline void
1033transpose_msg_vecs_sse2(const uint8_t *const *inputs, size_t block_offset, __m128i out_msg[16])
1034{
1035  size_t i;
1036
1037  out_msg[0]  = loadu_sse2(&inputs[0][block_offset + 0 * sizeof(__m128i)]);
1038  out_msg[1]  = loadu_sse2(&inputs[1][block_offset + 0 * sizeof(__m128i)]);
1039  out_msg[2]  = loadu_sse2(&inputs[2][block_offset + 0 * sizeof(__m128i)]);
1040  out_msg[3]  = loadu_sse2(&inputs[3][block_offset + 0 * sizeof(__m128i)]);
1041  out_msg[4]  = loadu_sse2(&inputs[0][block_offset + 1 * sizeof(__m128i)]);
1042  out_msg[5]  = loadu_sse2(&inputs[1][block_offset + 1 * sizeof(__m128i)]);
1043  out_msg[6]  = loadu_sse2(&inputs[2][block_offset + 1 * sizeof(__m128i)]);
1044  out_msg[7]  = loadu_sse2(&inputs[3][block_offset + 1 * sizeof(__m128i)]);
1045  out_msg[8]  = loadu_sse2(&inputs[0][block_offset + 2 * sizeof(__m128i)]);
1046  out_msg[9]  = loadu_sse2(&inputs[1][block_offset + 2 * sizeof(__m128i)]);
1047  out_msg[10] = loadu_sse2(&inputs[2][block_offset + 2 * sizeof(__m128i)]);
1048  out_msg[11] = loadu_sse2(&inputs[3][block_offset + 2 * sizeof(__m128i)]);
1049  out_msg[12] = loadu_sse2(&inputs[0][block_offset + 3 * sizeof(__m128i)]);
1050  out_msg[13] = loadu_sse2(&inputs[1][block_offset + 3 * sizeof(__m128i)]);
1051  out_msg[14] = loadu_sse2(&inputs[2][block_offset + 3 * sizeof(__m128i)]);
1052  out_msg[15] = loadu_sse2(&inputs[3][block_offset + 3 * sizeof(__m128i)]);
1053
1054  for (i = 0; i < 4; i++) {
1055    _mm_prefetch((const void *)&inputs[i][block_offset + 256], _MM_HINT_T0);
1056  }
1057  transpose_vecs_sse2(&out_msg[0]);
1058  transpose_vecs_sse2(&out_msg[4]);
1059  transpose_vecs_sse2(&out_msg[8]);
1060  transpose_vecs_sse2(&out_msg[12]);
1061}
1062
1063static inline void
1064load_counters_sse2(uint64_t counter, int increment_counter, __m128i *out_lo, __m128i *out_hi)
1065{
1066  const __m128i mask  = _mm_set1_epi32(-increment_counter);
1067  const __m128i add0  = _mm_set_epi32(3, 2, 1, 0);
1068  const __m128i add1  = _mm_and_si128(mask, add0);
1069  __m128i       l     = _mm_add_epi32(_mm_set1_epi32((int32_t)counter), add1);
1070  __m128i       carry = _mm_cmpgt_epi32(
1071      _mm_xor_si128(add1, _mm_set1_epi32(0x80000000)), _mm_xor_si128(l, _mm_set1_epi32(0x80000000))
1072  );
1073  __m128i h = _mm_sub_epi32(_mm_set1_epi32((int32_t)(counter >> 32)), carry);
1074
1075  *out_lo = l;
1076  *out_hi = h;
1077}
1078
1079void
1080blake3_hash4_sse2(
1081    const uint8_t *const *inputs,
1082    size_t                blocks,
1083    const uint32_t        key[8],
1084    uint64_t              counter,
1085    int                   increment_counter,
1086    uint8_t               flags,
1087    uint8_t               flags_start,
1088    uint8_t               flags_end,
1089    uint8_t              *out_bytes
1090)
1091{
1092  __m128i h_vecs[8];
1093  __m128i counter_low_vec, counter_high_vec;
1094  uint8_t block_flags;
1095  size_t  block;
1096
1097  h_vecs[0] = set1_sse2(key[0]);
1098  h_vecs[1] = set1_sse2(key[1]);
1099  h_vecs[2] = set1_sse2(key[2]);
1100  h_vecs[3] = set1_sse2(key[3]);
1101  h_vecs[4] = set1_sse2(key[4]);
1102  h_vecs[5] = set1_sse2(key[5]);
1103  h_vecs[6] = set1_sse2(key[6]);
1104  h_vecs[7] = set1_sse2(key[7]);
1105
1106  load_counters_sse2(counter, increment_counter, &counter_low_vec, &counter_high_vec);
1107  block_flags = flags | flags_start;
1108
1109  for (block = 0; block < blocks; block++) {
1110    __m128i block_len_vec;
1111    __m128i block_flags_vec;
1112    __m128i msg_vecs[16];
1113    __m128i v[16];
1114
1115    if (block + 1 == blocks) {
1116      block_flags |= flags_end;
1117    }
1118    block_len_vec   = set1_sse2(BLAKE3_BLOCK_LEN);
1119    block_flags_vec = set1_sse2(block_flags);
1120    transpose_msg_vecs_sse2(inputs, block * BLAKE3_BLOCK_LEN, msg_vecs);
1121
1122    v[0]  = h_vecs[0];
1123    v[1]  = h_vecs[1];
1124    v[2]  = h_vecs[2];
1125    v[3]  = h_vecs[3];
1126    v[4]  = h_vecs[4];
1127    v[5]  = h_vecs[5];
1128    v[6]  = h_vecs[6];
1129    v[7]  = h_vecs[7];
1130    v[8]  = set1_sse2(IV[0]);
1131    v[9]  = set1_sse2(IV[1]);
1132    v[10] = set1_sse2(IV[2]);
1133    v[11] = set1_sse2(IV[3]);
1134    v[12] = counter_low_vec;
1135    v[13] = counter_high_vec;
1136    v[14] = block_len_vec;
1137    v[15] = block_flags_vec;
1138
1139    round_fn_sse2(v, msg_vecs, 0);
1140    round_fn_sse2(v, msg_vecs, 1);
1141    round_fn_sse2(v, msg_vecs, 2);
1142    round_fn_sse2(v, msg_vecs, 3);
1143    round_fn_sse2(v, msg_vecs, 4);
1144    round_fn_sse2(v, msg_vecs, 5);
1145    round_fn_sse2(v, msg_vecs, 6);
1146
1147    h_vecs[0] = xorv_sse2(v[0], v[8]);
1148    h_vecs[1] = xorv_sse2(v[1], v[9]);
1149    h_vecs[2] = xorv_sse2(v[2], v[10]);
1150    h_vecs[3] = xorv_sse2(v[3], v[11]);
1151    h_vecs[4] = xorv_sse2(v[4], v[12]);
1152    h_vecs[5] = xorv_sse2(v[5], v[13]);
1153    h_vecs[6] = xorv_sse2(v[6], v[14]);
1154    h_vecs[7] = xorv_sse2(v[7], v[15]);
1155
1156    block_flags = flags;
1157  }
1158
1159  transpose_vecs_sse2(&h_vecs[0]);
1160  transpose_vecs_sse2(&h_vecs[4]);
1161  storeu_sse2(h_vecs[0], &out_bytes[0 * sizeof(__m128i)]);
1162  storeu_sse2(h_vecs[4], &out_bytes[1 * sizeof(__m128i)]);
1163  storeu_sse2(h_vecs[1], &out_bytes[2 * sizeof(__m128i)]);
1164  storeu_sse2(h_vecs[5], &out_bytes[3 * sizeof(__m128i)]);
1165  storeu_sse2(h_vecs[2], &out_bytes[4 * sizeof(__m128i)]);
1166  storeu_sse2(h_vecs[6], &out_bytes[5 * sizeof(__m128i)]);
1167  storeu_sse2(h_vecs[3], &out_bytes[6 * sizeof(__m128i)]);
1168  storeu_sse2(h_vecs[7], &out_bytes[7 * sizeof(__m128i)]);
1169}
1170
1171static inline void
1172hash_one_sse2(
1173    const uint8_t *input,
1174    size_t         blocks,
1175    const uint32_t key[8],
1176    uint64_t       counter,
1177    uint8_t        flags,
1178    uint8_t        flags_start,
1179    uint8_t        flags_end,
1180    uint8_t        out_bytes[BLAKE3_OUT_LEN]
1181)
1182{
1183  uint32_t cv[8];
1184  uint8_t  block_flags;
1185
1186  memcpy(cv, key, BLAKE3_KEY_LEN);
1187  block_flags = flags | flags_start;
1188  while (blocks > 0) {
1189    if (blocks == 1) {
1190      block_flags |= flags_end;
1191    }
1192    blake3_compress_in_place_sse2(cv, input, BLAKE3_BLOCK_LEN, counter, block_flags);
1193    input = &input[BLAKE3_BLOCK_LEN];
1194    blocks -= 1;
1195    block_flags = flags;
1196  }
1197  memcpy(out_bytes, cv, BLAKE3_OUT_LEN);
1198}
1199
1200void
1201blake3_hash_many_sse2(
1202    const uint8_t *const *inputs,
1203    size_t                num_inputs,
1204    size_t                blocks,
1205    const uint32_t        key[8],
1206    uint64_t              counter,
1207    int                   increment_counter,
1208    uint8_t               flags,
1209    uint8_t               flags_start,
1210    uint8_t               flags_end,
1211    uint8_t              *out_bytes
1212)
1213{
1214  while (num_inputs >= DEGREE_SSE2) {
1215    blake3_hash4_sse2(
1216        inputs, blocks, key, counter, increment_counter, flags, flags_start, flags_end, out_bytes
1217    );
1218    if (increment_counter) {
1219      counter += DEGREE_SSE2;
1220    }
1221    inputs += DEGREE_SSE2;
1222    num_inputs -= DEGREE_SSE2;
1223    out_bytes = &out_bytes[DEGREE_SSE2 * BLAKE3_OUT_LEN];
1224  }
1225  while (num_inputs > 0) {
1226    hash_one_sse2(inputs[0], blocks, key, counter, flags, flags_start, flags_end, out_bytes);
1227    if (increment_counter) {
1228      counter += 1;
1229    }
1230    inputs += 1;
1231    num_inputs -= 1;
1232    out_bytes = &out_bytes[BLAKE3_OUT_LEN];
1233  }
1234}
1235#if defined(__clang__)
1236#pragma clang attribute pop
1237#elif defined(__GNUC__)
1238#pragma GCC pop_options
1239#endif
1240
1241#if defined(__clang__)
1242#pragma clang attribute push(__attribute__((target("sse4.1"))), apply_to = function)
1243#elif defined(__GNUC__)
1244#pragma GCC push_options
1245#pragma GCC target("sse4.1")
1246#endif
1247#define DEGREE_SSE41 4
1248
1249#define _mm_shuffle_ps2_sse41(a, b, c)                                                             \
1250  (_mm_castps_si128(_mm_shuffle_ps(_mm_castsi128_ps(a), _mm_castsi128_ps(b), (c))))
1251
1252static inline __m128i
1253loadu_sse41(const uint8_t src[16])
1254{
1255  return _mm_loadu_si128((const __m128i *)src);
1256}
1257
1258static inline void
1259storeu_sse41(__m128i src, uint8_t dest[16])
1260{
1261  _mm_storeu_si128((__m128i *)dest, src);
1262}
1263
1264static inline __m128i
1265addv_sse41(__m128i a, __m128i b)
1266{
1267  return _mm_add_epi32(a, b);
1268}
1269
1270static inline __m128i
1271xorv_sse41(__m128i a, __m128i b)
1272{
1273  return _mm_xor_si128(a, b);
1274}
1275
1276static inline __m128i
1277set1_sse41(uint32_t x)
1278{
1279  return _mm_set1_epi32((int32_t)x);
1280}
1281
1282static inline __m128i
1283set4_sse41(uint32_t a, uint32_t b, uint32_t c, uint32_t d)
1284{
1285  return _mm_setr_epi32((int32_t)a, (int32_t)b, (int32_t)c, (int32_t)d);
1286}
1287
1288static inline __m128i
1289rot16_sse41(__m128i x)
1290{
1291  return _mm_shuffle_epi8(x, _mm_set_epi8(13, 12, 15, 14, 9, 8, 11, 10, 5, 4, 7, 6, 1, 0, 3, 2));
1292}
1293
1294static inline __m128i
1295rot12_sse41(__m128i x)
1296{
1297  return xorv_sse41(_mm_srli_epi32(x, 12), _mm_slli_epi32(x, 32 - 12));
1298}
1299
1300static inline __m128i
1301rot8_sse41(__m128i x)
1302{
1303  return _mm_shuffle_epi8(x, _mm_set_epi8(12, 15, 14, 13, 8, 11, 10, 9, 4, 7, 6, 5, 0, 3, 2, 1));
1304}
1305
1306static inline __m128i
1307rot7_sse41(__m128i x)
1308{
1309  return xorv_sse41(_mm_srli_epi32(x, 7), _mm_slli_epi32(x, 32 - 7));
1310}
1311
1312static inline void
1313g1_sse41(__m128i *row0, __m128i *row1, __m128i *row2, __m128i *row3, __m128i m)
1314{
1315  *row0 = addv_sse41(addv_sse41(*row0, m), *row1);
1316  *row3 = xorv_sse41(*row3, *row0);
1317  *row3 = rot16_sse41(*row3);
1318  *row2 = addv_sse41(*row2, *row3);
1319  *row1 = xorv_sse41(*row1, *row2);
1320  *row1 = rot12_sse41(*row1);
1321}
1322
1323static inline void
1324g2_sse41(__m128i *row0, __m128i *row1, __m128i *row2, __m128i *row3, __m128i m)
1325{
1326  *row0 = addv_sse41(addv_sse41(*row0, m), *row1);
1327  *row3 = xorv_sse41(*row3, *row0);
1328  *row3 = rot8_sse41(*row3);
1329  *row2 = addv_sse41(*row2, *row3);
1330  *row1 = xorv_sse41(*row1, *row2);
1331  *row1 = rot7_sse41(*row1);
1332}
1333
1334static inline void
1335diagonalize_sse41(__m128i *row0, __m128i *row2, __m128i *row3)
1336{
1337  *row0 = _mm_shuffle_epi32(*row0, _MM_SHUFFLE(2, 1, 0, 3));
1338  *row3 = _mm_shuffle_epi32(*row3, _MM_SHUFFLE(1, 0, 3, 2));
1339  *row2 = _mm_shuffle_epi32(*row2, _MM_SHUFFLE(0, 3, 2, 1));
1340}
1341
1342static inline void
1343undiagonalize_sse41(__m128i *row0, __m128i *row2, __m128i *row3)
1344{
1345  *row0 = _mm_shuffle_epi32(*row0, _MM_SHUFFLE(0, 3, 2, 1));
1346  *row3 = _mm_shuffle_epi32(*row3, _MM_SHUFFLE(1, 0, 3, 2));
1347  *row2 = _mm_shuffle_epi32(*row2, _MM_SHUFFLE(2, 1, 0, 3));
1348}
1349
1350static inline void
1351compress_pre_sse41(
1352    __m128i        rows[4],
1353    const uint32_t cv[8],
1354    const uint8_t  block[BLAKE3_BLOCK_LEN],
1355    uint8_t        block_len,
1356    uint64_t       counter,
1357    uint8_t        flags
1358)
1359{
1360  __m128i m0, m1, m2, m3;
1361  __m128i t0, t1, t2, t3, tt;
1362
1363  rows[0] = loadu_sse41((uint8_t *)&cv[0]);
1364  rows[1] = loadu_sse41((uint8_t *)&cv[4]);
1365  rows[2] = set4_sse41(IV[0], IV[1], IV[2], IV[3]);
1366  rows[3] =
1367      set4_sse41(counter_low(counter), counter_high(counter), (uint32_t)block_len, (uint32_t)flags);
1368
1369  m0 = loadu_sse41(&block[sizeof(__m128i) * 0]);
1370  m1 = loadu_sse41(&block[sizeof(__m128i) * 1]);
1371  m2 = loadu_sse41(&block[sizeof(__m128i) * 2]);
1372  m3 = loadu_sse41(&block[sizeof(__m128i) * 3]);
1373
1374  /* round 1 */
1375  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(2, 0, 2, 0));
1376  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1377  t1 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 3, 1));
1378  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1379  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1380  t2 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(2, 0, 2, 0));
1381  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(2, 1, 0, 3));
1382  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1383  t3 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 1, 3, 1));
1384  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(2, 1, 0, 3));
1385  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1386  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1387  m0 = t0;
1388  m1 = t1;
1389  m2 = t2;
1390  m3 = t3;
1391
1392  /* round 2 */
1393  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
1394  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
1395  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1396  t1 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
1397  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
1398  t1 = _mm_blend_epi16(tt, t1, 0xCC);
1399  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1400  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1401  t2 = _mm_unpacklo_epi64(m3, m1);
1402  tt = _mm_blend_epi16(t2, m2, 0xC0);
1403  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
1404  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1405  t3 = _mm_unpackhi_epi32(m1, m3);
1406  tt = _mm_unpacklo_epi32(m2, t3);
1407  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
1408  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1409  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1410  m0 = t0;
1411  m1 = t1;
1412  m2 = t2;
1413  m3 = t3;
1414
1415  /* round 3 */
1416  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
1417  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
1418  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1419  t1 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
1420  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
1421  t1 = _mm_blend_epi16(tt, t1, 0xCC);
1422  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1423  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1424  t2 = _mm_unpacklo_epi64(m3, m1);
1425  tt = _mm_blend_epi16(t2, m2, 0xC0);
1426  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
1427  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1428  t3 = _mm_unpackhi_epi32(m1, m3);
1429  tt = _mm_unpacklo_epi32(m2, t3);
1430  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
1431  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1432  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1433  m0 = t0;
1434  m1 = t1;
1435  m2 = t2;
1436  m3 = t3;
1437
1438  /* round 4 */
1439  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
1440  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
1441  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1442  t1 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
1443  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
1444  t1 = _mm_blend_epi16(tt, t1, 0xCC);
1445  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1446  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1447  t2 = _mm_unpacklo_epi64(m3, m1);
1448  tt = _mm_blend_epi16(t2, m2, 0xC0);
1449  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
1450  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1451  t3 = _mm_unpackhi_epi32(m1, m3);
1452  tt = _mm_unpacklo_epi32(m2, t3);
1453  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
1454  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1455  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1456  m0 = t0;
1457  m1 = t1;
1458  m2 = t2;
1459  m3 = t3;
1460
1461  /* round 5 */
1462  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
1463  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
1464  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1465  t1 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
1466  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
1467  t1 = _mm_blend_epi16(tt, t1, 0xCC);
1468  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1469  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1470  t2 = _mm_unpacklo_epi64(m3, m1);
1471  tt = _mm_blend_epi16(t2, m2, 0xC0);
1472  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
1473  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1474  t3 = _mm_unpackhi_epi32(m1, m3);
1475  tt = _mm_unpacklo_epi32(m2, t3);
1476  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
1477  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1478  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1479  m0 = t0;
1480  m1 = t1;
1481  m2 = t2;
1482  m3 = t3;
1483
1484  /* round 6 */
1485  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
1486  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
1487  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1488  t1 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
1489  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
1490  t1 = _mm_blend_epi16(tt, t1, 0xCC);
1491  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1492  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1493  t2 = _mm_unpacklo_epi64(m3, m1);
1494  tt = _mm_blend_epi16(t2, m2, 0xC0);
1495  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
1496  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1497  t3 = _mm_unpackhi_epi32(m1, m3);
1498  tt = _mm_unpacklo_epi32(m2, t3);
1499  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
1500  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1501  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1502  m0 = t0;
1503  m1 = t1;
1504  m2 = t2;
1505  m3 = t3;
1506
1507  /* round 7 */
1508  t0 = _mm_shuffle_ps2_sse41(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
1509  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
1510  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t0);
1511  t1 = _mm_shuffle_ps2_sse41(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
1512  tt = _mm_shuffle_epi32(m0, _MM_SHUFFLE(0, 0, 3, 3));
1513  t1 = _mm_blend_epi16(tt, t1, 0xCC);
1514  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t1);
1515  diagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1516  t2 = _mm_unpacklo_epi64(m3, m1);
1517  tt = _mm_blend_epi16(t2, m2, 0xC0);
1518  t2 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(1, 3, 2, 0));
1519  g1_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t2);
1520  t3 = _mm_unpackhi_epi32(m1, m3);
1521  tt = _mm_unpacklo_epi32(m2, t3);
1522  t3 = _mm_shuffle_epi32(tt, _MM_SHUFFLE(0, 1, 3, 2));
1523  g2_sse41(&rows[0], &rows[1], &rows[2], &rows[3], t3);
1524  undiagonalize_sse41(&rows[0], &rows[2], &rows[3]);
1525}
1526
1527void
1528blake3_compress_in_place_sse41(
1529    uint32_t      cv[8],
1530    const uint8_t block[BLAKE3_BLOCK_LEN],
1531    uint8_t       block_len,
1532    uint64_t      counter,
1533    uint8_t       flags
1534)
1535{
1536  __m128i rows[4];
1537
1538  compress_pre_sse41(rows, cv, block, block_len, counter, flags);
1539  storeu_sse41(xorv_sse41(rows[0], rows[2]), (uint8_t *)&cv[0]);
1540  storeu_sse41(xorv_sse41(rows[1], rows[3]), (uint8_t *)&cv[4]);
1541}
1542
1543void
1544blake3_compress_xof_sse41(
1545    const uint32_t cv[8],
1546    const uint8_t  block[BLAKE3_BLOCK_LEN],
1547    uint8_t        block_len,
1548    uint64_t       counter,
1549    uint8_t        flags,
1550    uint8_t        out_buf[64]
1551)
1552{
1553  __m128i rows[4];
1554
1555  compress_pre_sse41(rows, cv, block, block_len, counter, flags);
1556  storeu_sse41(xorv_sse41(rows[0], rows[2]), &out_buf[0]);
1557  storeu_sse41(xorv_sse41(rows[1], rows[3]), &out_buf[16]);
1558  storeu_sse41(xorv_sse41(rows[2], loadu_sse41((uint8_t *)&cv[0])), &out_buf[32]);
1559  storeu_sse41(xorv_sse41(rows[3], loadu_sse41((uint8_t *)&cv[4])), &out_buf[48]);
1560}
1561
1562static inline void
1563round_fn_sse41(__m128i v[16], __m128i m[16], size_t r)
1564{
1565  v[0]  = addv_sse41(v[0], m[(size_t)MSG_SCHEDULE[r][0]]);
1566  v[1]  = addv_sse41(v[1], m[(size_t)MSG_SCHEDULE[r][2]]);
1567  v[2]  = addv_sse41(v[2], m[(size_t)MSG_SCHEDULE[r][4]]);
1568  v[3]  = addv_sse41(v[3], m[(size_t)MSG_SCHEDULE[r][6]]);
1569  v[0]  = addv_sse41(v[0], v[4]);
1570  v[1]  = addv_sse41(v[1], v[5]);
1571  v[2]  = addv_sse41(v[2], v[6]);
1572  v[3]  = addv_sse41(v[3], v[7]);
1573  v[12] = xorv_sse41(v[12], v[0]);
1574  v[13] = xorv_sse41(v[13], v[1]);
1575  v[14] = xorv_sse41(v[14], v[2]);
1576  v[15] = xorv_sse41(v[15], v[3]);
1577  v[12] = rot16_sse41(v[12]);
1578  v[13] = rot16_sse41(v[13]);
1579  v[14] = rot16_sse41(v[14]);
1580  v[15] = rot16_sse41(v[15]);
1581  v[8]  = addv_sse41(v[8], v[12]);
1582  v[9]  = addv_sse41(v[9], v[13]);
1583  v[10] = addv_sse41(v[10], v[14]);
1584  v[11] = addv_sse41(v[11], v[15]);
1585  v[4]  = xorv_sse41(v[4], v[8]);
1586  v[5]  = xorv_sse41(v[5], v[9]);
1587  v[6]  = xorv_sse41(v[6], v[10]);
1588  v[7]  = xorv_sse41(v[7], v[11]);
1589  v[4]  = rot12_sse41(v[4]);
1590  v[5]  = rot12_sse41(v[5]);
1591  v[6]  = rot12_sse41(v[6]);
1592  v[7]  = rot12_sse41(v[7]);
1593  v[0]  = addv_sse41(v[0], m[(size_t)MSG_SCHEDULE[r][1]]);
1594  v[1]  = addv_sse41(v[1], m[(size_t)MSG_SCHEDULE[r][3]]);
1595  v[2]  = addv_sse41(v[2], m[(size_t)MSG_SCHEDULE[r][5]]);
1596  v[3]  = addv_sse41(v[3], m[(size_t)MSG_SCHEDULE[r][7]]);
1597  v[0]  = addv_sse41(v[0], v[4]);
1598  v[1]  = addv_sse41(v[1], v[5]);
1599  v[2]  = addv_sse41(v[2], v[6]);
1600  v[3]  = addv_sse41(v[3], v[7]);
1601  v[12] = xorv_sse41(v[12], v[0]);
1602  v[13] = xorv_sse41(v[13], v[1]);
1603  v[14] = xorv_sse41(v[14], v[2]);
1604  v[15] = xorv_sse41(v[15], v[3]);
1605  v[12] = rot8_sse41(v[12]);
1606  v[13] = rot8_sse41(v[13]);
1607  v[14] = rot8_sse41(v[14]);
1608  v[15] = rot8_sse41(v[15]);
1609  v[8]  = addv_sse41(v[8], v[12]);
1610  v[9]  = addv_sse41(v[9], v[13]);
1611  v[10] = addv_sse41(v[10], v[14]);
1612  v[11] = addv_sse41(v[11], v[15]);
1613  v[4]  = xorv_sse41(v[4], v[8]);
1614  v[5]  = xorv_sse41(v[5], v[9]);
1615  v[6]  = xorv_sse41(v[6], v[10]);
1616  v[7]  = xorv_sse41(v[7], v[11]);
1617  v[4]  = rot7_sse41(v[4]);
1618  v[5]  = rot7_sse41(v[5]);
1619  v[6]  = rot7_sse41(v[6]);
1620  v[7]  = rot7_sse41(v[7]);
1621
1622  v[0]  = addv_sse41(v[0], m[(size_t)MSG_SCHEDULE[r][8]]);
1623  v[1]  = addv_sse41(v[1], m[(size_t)MSG_SCHEDULE[r][10]]);
1624  v[2]  = addv_sse41(v[2], m[(size_t)MSG_SCHEDULE[r][12]]);
1625  v[3]  = addv_sse41(v[3], m[(size_t)MSG_SCHEDULE[r][14]]);
1626  v[0]  = addv_sse41(v[0], v[5]);
1627  v[1]  = addv_sse41(v[1], v[6]);
1628  v[2]  = addv_sse41(v[2], v[7]);
1629  v[3]  = addv_sse41(v[3], v[4]);
1630  v[15] = xorv_sse41(v[15], v[0]);
1631  v[12] = xorv_sse41(v[12], v[1]);
1632  v[13] = xorv_sse41(v[13], v[2]);
1633  v[14] = xorv_sse41(v[14], v[3]);
1634  v[15] = rot16_sse41(v[15]);
1635  v[12] = rot16_sse41(v[12]);
1636  v[13] = rot16_sse41(v[13]);
1637  v[14] = rot16_sse41(v[14]);
1638  v[10] = addv_sse41(v[10], v[15]);
1639  v[11] = addv_sse41(v[11], v[12]);
1640  v[8]  = addv_sse41(v[8], v[13]);
1641  v[9]  = addv_sse41(v[9], v[14]);
1642  v[5]  = xorv_sse41(v[5], v[10]);
1643  v[6]  = xorv_sse41(v[6], v[11]);
1644  v[7]  = xorv_sse41(v[7], v[8]);
1645  v[4]  = xorv_sse41(v[4], v[9]);
1646  v[5]  = rot12_sse41(v[5]);
1647  v[6]  = rot12_sse41(v[6]);
1648  v[7]  = rot12_sse41(v[7]);
1649  v[4]  = rot12_sse41(v[4]);
1650  v[0]  = addv_sse41(v[0], m[(size_t)MSG_SCHEDULE[r][9]]);
1651  v[1]  = addv_sse41(v[1], m[(size_t)MSG_SCHEDULE[r][11]]);
1652  v[2]  = addv_sse41(v[2], m[(size_t)MSG_SCHEDULE[r][13]]);
1653  v[3]  = addv_sse41(v[3], m[(size_t)MSG_SCHEDULE[r][15]]);
1654  v[0]  = addv_sse41(v[0], v[5]);
1655  v[1]  = addv_sse41(v[1], v[6]);
1656  v[2]  = addv_sse41(v[2], v[7]);
1657  v[3]  = addv_sse41(v[3], v[4]);
1658  v[15] = xorv_sse41(v[15], v[0]);
1659  v[12] = xorv_sse41(v[12], v[1]);
1660  v[13] = xorv_sse41(v[13], v[2]);
1661  v[14] = xorv_sse41(v[14], v[3]);
1662  v[15] = rot8_sse41(v[15]);
1663  v[12] = rot8_sse41(v[12]);
1664  v[13] = rot8_sse41(v[13]);
1665  v[14] = rot8_sse41(v[14]);
1666  v[10] = addv_sse41(v[10], v[15]);
1667  v[11] = addv_sse41(v[11], v[12]);
1668  v[8]  = addv_sse41(v[8], v[13]);
1669  v[9]  = addv_sse41(v[9], v[14]);
1670  v[5]  = xorv_sse41(v[5], v[10]);
1671  v[6]  = xorv_sse41(v[6], v[11]);
1672  v[7]  = xorv_sse41(v[7], v[8]);
1673  v[4]  = xorv_sse41(v[4], v[9]);
1674  v[5]  = rot7_sse41(v[5]);
1675  v[6]  = rot7_sse41(v[6]);
1676  v[7]  = rot7_sse41(v[7]);
1677  v[4]  = rot7_sse41(v[4]);
1678}
1679
1680static inline void
1681transpose_vecs_sse41(__m128i vecs[DEGREE_SSE41])
1682{
1683  __m128i ab_01 = _mm_unpacklo_epi32(vecs[0], vecs[1]);
1684  __m128i ab_23 = _mm_unpackhi_epi32(vecs[0], vecs[1]);
1685  __m128i cd_01 = _mm_unpacklo_epi32(vecs[2], vecs[3]);
1686  __m128i cd_23 = _mm_unpackhi_epi32(vecs[2], vecs[3]);
1687
1688  __m128i abcd_0 = _mm_unpacklo_epi64(ab_01, cd_01);
1689  __m128i abcd_1 = _mm_unpackhi_epi64(ab_01, cd_01);
1690  __m128i abcd_2 = _mm_unpacklo_epi64(ab_23, cd_23);
1691  __m128i abcd_3 = _mm_unpackhi_epi64(ab_23, cd_23);
1692
1693  vecs[0] = abcd_0;
1694  vecs[1] = abcd_1;
1695  vecs[2] = abcd_2;
1696  vecs[3] = abcd_3;
1697}
1698
1699static inline void
1700transpose_msg_vecs_sse41(const uint8_t *const *inputs, size_t block_offset, __m128i out_msg[16])
1701{
1702  size_t i;
1703
1704  out_msg[0]  = loadu_sse41(&inputs[0][block_offset + 0 * sizeof(__m128i)]);
1705  out_msg[1]  = loadu_sse41(&inputs[1][block_offset + 0 * sizeof(__m128i)]);
1706  out_msg[2]  = loadu_sse41(&inputs[2][block_offset + 0 * sizeof(__m128i)]);
1707  out_msg[3]  = loadu_sse41(&inputs[3][block_offset + 0 * sizeof(__m128i)]);
1708  out_msg[4]  = loadu_sse41(&inputs[0][block_offset + 1 * sizeof(__m128i)]);
1709  out_msg[5]  = loadu_sse41(&inputs[1][block_offset + 1 * sizeof(__m128i)]);
1710  out_msg[6]  = loadu_sse41(&inputs[2][block_offset + 1 * sizeof(__m128i)]);
1711  out_msg[7]  = loadu_sse41(&inputs[3][block_offset + 1 * sizeof(__m128i)]);
1712  out_msg[8]  = loadu_sse41(&inputs[0][block_offset + 2 * sizeof(__m128i)]);
1713  out_msg[9]  = loadu_sse41(&inputs[1][block_offset + 2 * sizeof(__m128i)]);
1714  out_msg[10] = loadu_sse41(&inputs[2][block_offset + 2 * sizeof(__m128i)]);
1715  out_msg[11] = loadu_sse41(&inputs[3][block_offset + 2 * sizeof(__m128i)]);
1716  out_msg[12] = loadu_sse41(&inputs[0][block_offset + 3 * sizeof(__m128i)]);
1717  out_msg[13] = loadu_sse41(&inputs[1][block_offset + 3 * sizeof(__m128i)]);
1718  out_msg[14] = loadu_sse41(&inputs[2][block_offset + 3 * sizeof(__m128i)]);
1719  out_msg[15] = loadu_sse41(&inputs[3][block_offset + 3 * sizeof(__m128i)]);
1720
1721  for (i = 0; i < 4; i++) {
1722    _mm_prefetch((const void *)&inputs[i][block_offset + 256], _MM_HINT_T0);
1723  }
1724  transpose_vecs_sse41(&out_msg[0]);
1725  transpose_vecs_sse41(&out_msg[4]);
1726  transpose_vecs_sse41(&out_msg[8]);
1727  transpose_vecs_sse41(&out_msg[12]);
1728}
1729
1730static inline void
1731load_counters_sse41(uint64_t counter, int increment_counter, __m128i *out_lo, __m128i *out_hi)
1732{
1733  const __m128i mask  = _mm_set1_epi32(-increment_counter);
1734  const __m128i add0  = _mm_set_epi32(3, 2, 1, 0);
1735  const __m128i add1  = _mm_and_si128(mask, add0);
1736  __m128i       l     = _mm_add_epi32(_mm_set1_epi32((int32_t)counter), add1);
1737  __m128i       carry = _mm_cmpgt_epi32(
1738      _mm_xor_si128(add1, _mm_set1_epi32(0x80000000)), _mm_xor_si128(l, _mm_set1_epi32(0x80000000))
1739  );
1740  __m128i h = _mm_sub_epi32(_mm_set1_epi32((int32_t)(counter >> 32)), carry);
1741
1742  *out_lo = l;
1743  *out_hi = h;
1744}
1745
1746void
1747blake3_hash4_sse41(
1748    const uint8_t *const *inputs,
1749    size_t                blocks,
1750    const uint32_t        key[8],
1751    uint64_t              counter,
1752    int                   increment_counter,
1753    uint8_t               flags,
1754    uint8_t               flags_start,
1755    uint8_t               flags_end,
1756    uint8_t              *out_bytes
1757)
1758{
1759  __m128i h_vecs[8];
1760  __m128i counter_low_vec, counter_high_vec;
1761  uint8_t block_flags;
1762  size_t  block;
1763
1764  h_vecs[0] = set1_sse41(key[0]);
1765  h_vecs[1] = set1_sse41(key[1]);
1766  h_vecs[2] = set1_sse41(key[2]);
1767  h_vecs[3] = set1_sse41(key[3]);
1768  h_vecs[4] = set1_sse41(key[4]);
1769  h_vecs[5] = set1_sse41(key[5]);
1770  h_vecs[6] = set1_sse41(key[6]);
1771  h_vecs[7] = set1_sse41(key[7]);
1772
1773  load_counters_sse41(counter, increment_counter, &counter_low_vec, &counter_high_vec);
1774  block_flags = flags | flags_start;
1775
1776  for (block = 0; block < blocks; block++) {
1777    __m128i block_len_vec;
1778    __m128i block_flags_vec;
1779    __m128i msg_vecs[16];
1780    __m128i v[16];
1781
1782    if (block + 1 == blocks) {
1783      block_flags |= flags_end;
1784    }
1785    block_len_vec   = set1_sse41(BLAKE3_BLOCK_LEN);
1786    block_flags_vec = set1_sse41(block_flags);
1787    transpose_msg_vecs_sse41(inputs, block * BLAKE3_BLOCK_LEN, msg_vecs);
1788
1789    v[0]  = h_vecs[0];
1790    v[1]  = h_vecs[1];
1791    v[2]  = h_vecs[2];
1792    v[3]  = h_vecs[3];
1793    v[4]  = h_vecs[4];
1794    v[5]  = h_vecs[5];
1795    v[6]  = h_vecs[6];
1796    v[7]  = h_vecs[7];
1797    v[8]  = set1_sse41(IV[0]);
1798    v[9]  = set1_sse41(IV[1]);
1799    v[10] = set1_sse41(IV[2]);
1800    v[11] = set1_sse41(IV[3]);
1801    v[12] = counter_low_vec;
1802    v[13] = counter_high_vec;
1803    v[14] = block_len_vec;
1804    v[15] = block_flags_vec;
1805
1806    round_fn_sse41(v, msg_vecs, 0);
1807    round_fn_sse41(v, msg_vecs, 1);
1808    round_fn_sse41(v, msg_vecs, 2);
1809    round_fn_sse41(v, msg_vecs, 3);
1810    round_fn_sse41(v, msg_vecs, 4);
1811    round_fn_sse41(v, msg_vecs, 5);
1812    round_fn_sse41(v, msg_vecs, 6);
1813
1814    h_vecs[0] = xorv_sse41(v[0], v[8]);
1815    h_vecs[1] = xorv_sse41(v[1], v[9]);
1816    h_vecs[2] = xorv_sse41(v[2], v[10]);
1817    h_vecs[3] = xorv_sse41(v[3], v[11]);
1818    h_vecs[4] = xorv_sse41(v[4], v[12]);
1819    h_vecs[5] = xorv_sse41(v[5], v[13]);
1820    h_vecs[6] = xorv_sse41(v[6], v[14]);
1821    h_vecs[7] = xorv_sse41(v[7], v[15]);
1822
1823    block_flags = flags;
1824  }
1825
1826  transpose_vecs_sse41(&h_vecs[0]);
1827  transpose_vecs_sse41(&h_vecs[4]);
1828  storeu_sse41(h_vecs[0], &out_bytes[0 * sizeof(__m128i)]);
1829  storeu_sse41(h_vecs[4], &out_bytes[1 * sizeof(__m128i)]);
1830  storeu_sse41(h_vecs[1], &out_bytes[2 * sizeof(__m128i)]);
1831  storeu_sse41(h_vecs[5], &out_bytes[3 * sizeof(__m128i)]);
1832  storeu_sse41(h_vecs[2], &out_bytes[4 * sizeof(__m128i)]);
1833  storeu_sse41(h_vecs[6], &out_bytes[5 * sizeof(__m128i)]);
1834  storeu_sse41(h_vecs[3], &out_bytes[6 * sizeof(__m128i)]);
1835  storeu_sse41(h_vecs[7], &out_bytes[7 * sizeof(__m128i)]);
1836}
1837
1838static inline void
1839hash_one_sse41(
1840    const uint8_t *input,
1841    size_t         blocks,
1842    const uint32_t key[8],
1843    uint64_t       counter,
1844    uint8_t        flags,
1845    uint8_t        flags_start,
1846    uint8_t        flags_end,
1847    uint8_t        out_bytes[BLAKE3_OUT_LEN]
1848)
1849{
1850  uint32_t cv[8];
1851  uint8_t  block_flags;
1852
1853  memcpy(cv, key, BLAKE3_KEY_LEN);
1854  block_flags = flags | flags_start;
1855  while (blocks > 0) {
1856    if (blocks == 1) {
1857      block_flags |= flags_end;
1858    }
1859    blake3_compress_in_place_sse41(cv, input, BLAKE3_BLOCK_LEN, counter, block_flags);
1860    input = &input[BLAKE3_BLOCK_LEN];
1861    blocks -= 1;
1862    block_flags = flags;
1863  }
1864  memcpy(out_bytes, cv, BLAKE3_OUT_LEN);
1865}
1866
1867void
1868blake3_hash_many_sse41(
1869    const uint8_t *const *inputs,
1870    size_t                num_inputs,
1871    size_t                blocks,
1872    const uint32_t        key[8],
1873    uint64_t              counter,
1874    int                   increment_counter,
1875    uint8_t               flags,
1876    uint8_t               flags_start,
1877    uint8_t               flags_end,
1878    uint8_t              *out_bytes
1879)
1880{
1881  while (num_inputs >= DEGREE_SSE41) {
1882    blake3_hash4_sse41(
1883        inputs, blocks, key, counter, increment_counter, flags, flags_start, flags_end, out_bytes
1884    );
1885    if (increment_counter) {
1886      counter += DEGREE_SSE41;
1887    }
1888    inputs += DEGREE_SSE41;
1889    num_inputs -= DEGREE_SSE41;
1890    out_bytes = &out_bytes[DEGREE_SSE41 * BLAKE3_OUT_LEN];
1891  }
1892  while (num_inputs > 0) {
1893    hash_one_sse41(inputs[0], blocks, key, counter, flags, flags_start, flags_end, out_bytes);
1894    if (increment_counter) {
1895      counter += 1;
1896    }
1897    inputs += 1;
1898    num_inputs -= 1;
1899    out_bytes = &out_bytes[BLAKE3_OUT_LEN];
1900  }
1901}
1902#if defined(__clang__)
1903#pragma clang attribute pop
1904#elif defined(__GNUC__)
1905#pragma GCC pop_options
1906#endif
1907
1908#if defined(__clang__)
1909#pragma clang attribute push(__attribute__((target("avx2"))), apply_to = function)
1910#elif defined(__GNUC__)
1911#pragma GCC push_options
1912#pragma GCC target("avx2")
1913#endif
1914#define DEGREE_AVX2 8
1915
1916static inline __m256i
1917loadu_avx2(const uint8_t src[32])
1918{
1919  return _mm256_loadu_si256((const __m256i *)src);
1920}
1921
1922static inline void
1923storeu_avx2(__m256i src, uint8_t dest[32])
1924{
1925  _mm256_storeu_si256((__m256i *)dest, src);
1926}
1927
1928static inline __m256i
1929addv_avx2(__m256i a, __m256i b)
1930{
1931  return _mm256_add_epi32(a, b);
1932}
1933
1934static inline __m256i
1935xorv_avx2(__m256i a, __m256i b)
1936{
1937  return _mm256_xor_si256(a, b);
1938}
1939
1940static inline __m256i
1941set1_avx2(uint32_t x)
1942{
1943  return _mm256_set1_epi32((int32_t)x);
1944}
1945
1946static inline __m256i
1947rot16_avx2(__m256i x)
1948{
1949  return _mm256_shuffle_epi8(
1950      x,
1951      _mm256_set_epi8(
1952          13,
1953          12,
1954          15,
1955          14,
1956          9,
1957          8,
1958          11,
1959          10,
1960          5,
1961          4,
1962          7,
1963          6,
1964          1,
1965          0,
1966          3,
1967          2,
1968          13,
1969          12,
1970          15,
1971          14,
1972          9,
1973          8,
1974          11,
1975          10,
1976          5,
1977          4,
1978          7,
1979          6,
1980          1,
1981          0,
1982          3,
1983          2
1984      )
1985  );
1986}
1987
1988static inline __m256i
1989rot12_avx2(__m256i x)
1990{
1991  return _mm256_or_si256(_mm256_srli_epi32(x, 12), _mm256_slli_epi32(x, 32 - 12));
1992}
1993
1994static inline __m256i
1995rot8_avx2(__m256i x)
1996{
1997  return _mm256_shuffle_epi8(
1998      x,
1999      _mm256_set_epi8(
2000          12,
2001          15,
2002          14,
2003          13,
2004          8,
2005          11,
2006          10,
2007          9,
2008          4,
2009          7,
2010          6,
2011          5,
2012          0,
2013          3,
2014          2,
2015          1,
2016          12,
2017          15,
2018          14,
2019          13,
2020          8,
2021          11,
2022          10,
2023          9,
2024          4,
2025          7,
2026          6,
2027          5,
2028          0,
2029          3,
2030          2,
2031          1
2032      )
2033  );
2034}
2035
2036static inline __m256i
2037rot7_avx2(__m256i x)
2038{
2039  return _mm256_or_si256(_mm256_srli_epi32(x, 7), _mm256_slli_epi32(x, 32 - 7));
2040}
2041
2042static inline void
2043round_fn_avx2(__m256i v[16], __m256i m[16], size_t r)
2044{
2045  v[0]  = addv_avx2(v[0], m[(size_t)MSG_SCHEDULE[r][0]]);
2046  v[1]  = addv_avx2(v[1], m[(size_t)MSG_SCHEDULE[r][2]]);
2047  v[2]  = addv_avx2(v[2], m[(size_t)MSG_SCHEDULE[r][4]]);
2048  v[3]  = addv_avx2(v[3], m[(size_t)MSG_SCHEDULE[r][6]]);
2049  v[0]  = addv_avx2(v[0], v[4]);
2050  v[1]  = addv_avx2(v[1], v[5]);
2051  v[2]  = addv_avx2(v[2], v[6]);
2052  v[3]  = addv_avx2(v[3], v[7]);
2053  v[12] = xorv_avx2(v[12], v[0]);
2054  v[13] = xorv_avx2(v[13], v[1]);
2055  v[14] = xorv_avx2(v[14], v[2]);
2056  v[15] = xorv_avx2(v[15], v[3]);
2057  v[12] = rot16_avx2(v[12]);
2058  v[13] = rot16_avx2(v[13]);
2059  v[14] = rot16_avx2(v[14]);
2060  v[15] = rot16_avx2(v[15]);
2061  v[8]  = addv_avx2(v[8], v[12]);
2062  v[9]  = addv_avx2(v[9], v[13]);
2063  v[10] = addv_avx2(v[10], v[14]);
2064  v[11] = addv_avx2(v[11], v[15]);
2065  v[4]  = xorv_avx2(v[4], v[8]);
2066  v[5]  = xorv_avx2(v[5], v[9]);
2067  v[6]  = xorv_avx2(v[6], v[10]);
2068  v[7]  = xorv_avx2(v[7], v[11]);
2069  v[4]  = rot12_avx2(v[4]);
2070  v[5]  = rot12_avx2(v[5]);
2071  v[6]  = rot12_avx2(v[6]);
2072  v[7]  = rot12_avx2(v[7]);
2073  v[0]  = addv_avx2(v[0], m[(size_t)MSG_SCHEDULE[r][1]]);
2074  v[1]  = addv_avx2(v[1], m[(size_t)MSG_SCHEDULE[r][3]]);
2075  v[2]  = addv_avx2(v[2], m[(size_t)MSG_SCHEDULE[r][5]]);
2076  v[3]  = addv_avx2(v[3], m[(size_t)MSG_SCHEDULE[r][7]]);
2077  v[0]  = addv_avx2(v[0], v[4]);
2078  v[1]  = addv_avx2(v[1], v[5]);
2079  v[2]  = addv_avx2(v[2], v[6]);
2080  v[3]  = addv_avx2(v[3], v[7]);
2081  v[12] = xorv_avx2(v[12], v[0]);
2082  v[13] = xorv_avx2(v[13], v[1]);
2083  v[14] = xorv_avx2(v[14], v[2]);
2084  v[15] = xorv_avx2(v[15], v[3]);
2085  v[12] = rot8_avx2(v[12]);
2086  v[13] = rot8_avx2(v[13]);
2087  v[14] = rot8_avx2(v[14]);
2088  v[15] = rot8_avx2(v[15]);
2089  v[8]  = addv_avx2(v[8], v[12]);
2090  v[9]  = addv_avx2(v[9], v[13]);
2091  v[10] = addv_avx2(v[10], v[14]);
2092  v[11] = addv_avx2(v[11], v[15]);
2093  v[4]  = xorv_avx2(v[4], v[8]);
2094  v[5]  = xorv_avx2(v[5], v[9]);
2095  v[6]  = xorv_avx2(v[6], v[10]);
2096  v[7]  = xorv_avx2(v[7], v[11]);
2097  v[4]  = rot7_avx2(v[4]);
2098  v[5]  = rot7_avx2(v[5]);
2099  v[6]  = rot7_avx2(v[6]);
2100  v[7]  = rot7_avx2(v[7]);
2101
2102  v[0]  = addv_avx2(v[0], m[(size_t)MSG_SCHEDULE[r][8]]);
2103  v[1]  = addv_avx2(v[1], m[(size_t)MSG_SCHEDULE[r][10]]);
2104  v[2]  = addv_avx2(v[2], m[(size_t)MSG_SCHEDULE[r][12]]);
2105  v[3]  = addv_avx2(v[3], m[(size_t)MSG_SCHEDULE[r][14]]);
2106  v[0]  = addv_avx2(v[0], v[5]);
2107  v[1]  = addv_avx2(v[1], v[6]);
2108  v[2]  = addv_avx2(v[2], v[7]);
2109  v[3]  = addv_avx2(v[3], v[4]);
2110  v[15] = xorv_avx2(v[15], v[0]);
2111  v[12] = xorv_avx2(v[12], v[1]);
2112  v[13] = xorv_avx2(v[13], v[2]);
2113  v[14] = xorv_avx2(v[14], v[3]);
2114  v[15] = rot16_avx2(v[15]);
2115  v[12] = rot16_avx2(v[12]);
2116  v[13] = rot16_avx2(v[13]);
2117  v[14] = rot16_avx2(v[14]);
2118  v[10] = addv_avx2(v[10], v[15]);
2119  v[11] = addv_avx2(v[11], v[12]);
2120  v[8]  = addv_avx2(v[8], v[13]);
2121  v[9]  = addv_avx2(v[9], v[14]);
2122  v[5]  = xorv_avx2(v[5], v[10]);
2123  v[6]  = xorv_avx2(v[6], v[11]);
2124  v[7]  = xorv_avx2(v[7], v[8]);
2125  v[4]  = xorv_avx2(v[4], v[9]);
2126  v[5]  = rot12_avx2(v[5]);
2127  v[6]  = rot12_avx2(v[6]);
2128  v[7]  = rot12_avx2(v[7]);
2129  v[4]  = rot12_avx2(v[4]);
2130  v[0]  = addv_avx2(v[0], m[(size_t)MSG_SCHEDULE[r][9]]);
2131  v[1]  = addv_avx2(v[1], m[(size_t)MSG_SCHEDULE[r][11]]);
2132  v[2]  = addv_avx2(v[2], m[(size_t)MSG_SCHEDULE[r][13]]);
2133  v[3]  = addv_avx2(v[3], m[(size_t)MSG_SCHEDULE[r][15]]);
2134  v[0]  = addv_avx2(v[0], v[5]);
2135  v[1]  = addv_avx2(v[1], v[6]);
2136  v[2]  = addv_avx2(v[2], v[7]);
2137  v[3]  = addv_avx2(v[3], v[4]);
2138  v[15] = xorv_avx2(v[15], v[0]);
2139  v[12] = xorv_avx2(v[12], v[1]);
2140  v[13] = xorv_avx2(v[13], v[2]);
2141  v[14] = xorv_avx2(v[14], v[3]);
2142  v[15] = rot8_avx2(v[15]);
2143  v[12] = rot8_avx2(v[12]);
2144  v[13] = rot8_avx2(v[13]);
2145  v[14] = rot8_avx2(v[14]);
2146  v[10] = addv_avx2(v[10], v[15]);
2147  v[11] = addv_avx2(v[11], v[12]);
2148  v[8]  = addv_avx2(v[8], v[13]);
2149  v[9]  = addv_avx2(v[9], v[14]);
2150  v[5]  = xorv_avx2(v[5], v[10]);
2151  v[6]  = xorv_avx2(v[6], v[11]);
2152  v[7]  = xorv_avx2(v[7], v[8]);
2153  v[4]  = xorv_avx2(v[4], v[9]);
2154  v[5]  = rot7_avx2(v[5]);
2155  v[6]  = rot7_avx2(v[6]);
2156  v[7]  = rot7_avx2(v[7]);
2157  v[4]  = rot7_avx2(v[4]);
2158}
2159
2160static inline void
2161transpose_vecs_avx2(__m256i vecs[DEGREE_AVX2])
2162{
2163  __m256i ab_0145 = _mm256_unpacklo_epi32(vecs[0], vecs[1]);
2164  __m256i ab_2367 = _mm256_unpackhi_epi32(vecs[0], vecs[1]);
2165  __m256i cd_0145 = _mm256_unpacklo_epi32(vecs[2], vecs[3]);
2166  __m256i cd_2367 = _mm256_unpackhi_epi32(vecs[2], vecs[3]);
2167  __m256i ef_0145 = _mm256_unpacklo_epi32(vecs[4], vecs[5]);
2168  __m256i ef_2367 = _mm256_unpackhi_epi32(vecs[4], vecs[5]);
2169  __m256i gh_0145 = _mm256_unpacklo_epi32(vecs[6], vecs[7]);
2170  __m256i gh_2367 = _mm256_unpackhi_epi32(vecs[6], vecs[7]);
2171
2172  __m256i abcd_04 = _mm256_unpacklo_epi64(ab_0145, cd_0145);
2173  __m256i abcd_15 = _mm256_unpackhi_epi64(ab_0145, cd_0145);
2174  __m256i abcd_26 = _mm256_unpacklo_epi64(ab_2367, cd_2367);
2175  __m256i abcd_37 = _mm256_unpackhi_epi64(ab_2367, cd_2367);
2176  __m256i efgh_04 = _mm256_unpacklo_epi64(ef_0145, gh_0145);
2177  __m256i efgh_15 = _mm256_unpackhi_epi64(ef_0145, gh_0145);
2178  __m256i efgh_26 = _mm256_unpacklo_epi64(ef_2367, gh_2367);
2179  __m256i efgh_37 = _mm256_unpackhi_epi64(ef_2367, gh_2367);
2180
2181  vecs[0] = _mm256_permute2x128_si256(abcd_04, efgh_04, 0x20);
2182  vecs[1] = _mm256_permute2x128_si256(abcd_15, efgh_15, 0x20);
2183  vecs[2] = _mm256_permute2x128_si256(abcd_26, efgh_26, 0x20);
2184  vecs[3] = _mm256_permute2x128_si256(abcd_37, efgh_37, 0x20);
2185  vecs[4] = _mm256_permute2x128_si256(abcd_04, efgh_04, 0x31);
2186  vecs[5] = _mm256_permute2x128_si256(abcd_15, efgh_15, 0x31);
2187  vecs[6] = _mm256_permute2x128_si256(abcd_26, efgh_26, 0x31);
2188  vecs[7] = _mm256_permute2x128_si256(abcd_37, efgh_37, 0x31);
2189}
2190
2191static inline void
2192transpose_msg_vecs_avx2(const uint8_t *const *inputs, size_t block_offset, __m256i out_msg[16])
2193{
2194  size_t i;
2195
2196  out_msg[0]  = loadu_avx2(&inputs[0][block_offset + 0 * sizeof(__m256i)]);
2197  out_msg[1]  = loadu_avx2(&inputs[1][block_offset + 0 * sizeof(__m256i)]);
2198  out_msg[2]  = loadu_avx2(&inputs[2][block_offset + 0 * sizeof(__m256i)]);
2199  out_msg[3]  = loadu_avx2(&inputs[3][block_offset + 0 * sizeof(__m256i)]);
2200  out_msg[4]  = loadu_avx2(&inputs[4][block_offset + 0 * sizeof(__m256i)]);
2201  out_msg[5]  = loadu_avx2(&inputs[5][block_offset + 0 * sizeof(__m256i)]);
2202  out_msg[6]  = loadu_avx2(&inputs[6][block_offset + 0 * sizeof(__m256i)]);
2203  out_msg[7]  = loadu_avx2(&inputs[7][block_offset + 0 * sizeof(__m256i)]);
2204  out_msg[8]  = loadu_avx2(&inputs[0][block_offset + 1 * sizeof(__m256i)]);
2205  out_msg[9]  = loadu_avx2(&inputs[1][block_offset + 1 * sizeof(__m256i)]);
2206  out_msg[10] = loadu_avx2(&inputs[2][block_offset + 1 * sizeof(__m256i)]);
2207  out_msg[11] = loadu_avx2(&inputs[3][block_offset + 1 * sizeof(__m256i)]);
2208  out_msg[12] = loadu_avx2(&inputs[4][block_offset + 1 * sizeof(__m256i)]);
2209  out_msg[13] = loadu_avx2(&inputs[5][block_offset + 1 * sizeof(__m256i)]);
2210  out_msg[14] = loadu_avx2(&inputs[6][block_offset + 1 * sizeof(__m256i)]);
2211  out_msg[15] = loadu_avx2(&inputs[7][block_offset + 1 * sizeof(__m256i)]);
2212
2213  for (i = 0; i < 8; i++) {
2214    _mm_prefetch((const void *)&inputs[i][block_offset + 256], _MM_HINT_T0);
2215  }
2216  transpose_vecs_avx2(&out_msg[0]);
2217  transpose_vecs_avx2(&out_msg[8]);
2218}
2219
2220static inline void
2221load_counters_avx2(uint64_t counter, int increment_counter, __m256i *out_lo, __m256i *out_hi)
2222{
2223  const __m256i mask  = _mm256_set1_epi32(-increment_counter);
2224  const __m256i add0  = _mm256_set_epi32(7, 6, 5, 4, 3, 2, 1, 0);
2225  const __m256i add1  = _mm256_and_si256(mask, add0);
2226  __m256i       l     = _mm256_add_epi32(_mm256_set1_epi32((int32_t)counter), add1);
2227  __m256i       carry = _mm256_xor_si256(add1, _mm256_set1_epi32(0x80000000));
2228  __m256i       comp  = _mm256_xor_si256(l, _mm256_set1_epi32(0x80000000));
2229  __m256i       gt    = _mm256_cmpgt_epi32(carry, comp);
2230  __m256i       h     = _mm256_sub_epi32(_mm256_set1_epi32((int32_t)(counter >> 32)), gt);
2231
2232  *out_lo = l;
2233  *out_hi = h;
2234}
2235
2236void
2237blake3_hash8_avx2(
2238    const uint8_t *const *inputs,
2239    size_t                blocks,
2240    const uint32_t        key[8],
2241    uint64_t              counter,
2242    int                   increment_counter,
2243    uint8_t               flags,
2244    uint8_t               flags_start,
2245    uint8_t               flags_end,
2246    uint8_t              *out_bytes
2247)
2248{
2249  __m256i h_vecs[8];
2250  __m256i counter_low_vec, counter_high_vec;
2251  uint8_t block_flags;
2252  size_t  block;
2253
2254  h_vecs[0] = set1_avx2(key[0]);
2255  h_vecs[1] = set1_avx2(key[1]);
2256  h_vecs[2] = set1_avx2(key[2]);
2257  h_vecs[3] = set1_avx2(key[3]);
2258  h_vecs[4] = set1_avx2(key[4]);
2259  h_vecs[5] = set1_avx2(key[5]);
2260  h_vecs[6] = set1_avx2(key[6]);
2261  h_vecs[7] = set1_avx2(key[7]);
2262
2263  load_counters_avx2(counter, increment_counter, &counter_low_vec, &counter_high_vec);
2264  block_flags = flags | flags_start;
2265
2266  for (block = 0; block < blocks; block++) {
2267    __m256i block_len_vec;
2268    __m256i block_flags_vec;
2269    __m256i msg_vecs[16];
2270    __m256i v[16];
2271
2272    if (block + 1 == blocks) {
2273      block_flags |= flags_end;
2274    }
2275    block_len_vec   = set1_avx2(BLAKE3_BLOCK_LEN);
2276    block_flags_vec = set1_avx2(block_flags);
2277    transpose_msg_vecs_avx2(inputs, block * BLAKE3_BLOCK_LEN, msg_vecs);
2278
2279    v[0]  = h_vecs[0];
2280    v[1]  = h_vecs[1];
2281    v[2]  = h_vecs[2];
2282    v[3]  = h_vecs[3];
2283    v[4]  = h_vecs[4];
2284    v[5]  = h_vecs[5];
2285    v[6]  = h_vecs[6];
2286    v[7]  = h_vecs[7];
2287    v[8]  = set1_avx2(IV[0]);
2288    v[9]  = set1_avx2(IV[1]);
2289    v[10] = set1_avx2(IV[2]);
2290    v[11] = set1_avx2(IV[3]);
2291    v[12] = counter_low_vec;
2292    v[13] = counter_high_vec;
2293    v[14] = block_len_vec;
2294    v[15] = block_flags_vec;
2295
2296    round_fn_avx2(v, msg_vecs, 0);
2297    round_fn_avx2(v, msg_vecs, 1);
2298    round_fn_avx2(v, msg_vecs, 2);
2299    round_fn_avx2(v, msg_vecs, 3);
2300    round_fn_avx2(v, msg_vecs, 4);
2301    round_fn_avx2(v, msg_vecs, 5);
2302    round_fn_avx2(v, msg_vecs, 6);
2303
2304    h_vecs[0] = xorv_avx2(v[0], v[8]);
2305    h_vecs[1] = xorv_avx2(v[1], v[9]);
2306    h_vecs[2] = xorv_avx2(v[2], v[10]);
2307    h_vecs[3] = xorv_avx2(v[3], v[11]);
2308    h_vecs[4] = xorv_avx2(v[4], v[12]);
2309    h_vecs[5] = xorv_avx2(v[5], v[13]);
2310    h_vecs[6] = xorv_avx2(v[6], v[14]);
2311    h_vecs[7] = xorv_avx2(v[7], v[15]);
2312
2313    block_flags = flags;
2314  }
2315
2316  transpose_vecs_avx2(h_vecs);
2317  storeu_avx2(h_vecs[0], &out_bytes[0 * sizeof(__m256i)]);
2318  storeu_avx2(h_vecs[1], &out_bytes[1 * sizeof(__m256i)]);
2319  storeu_avx2(h_vecs[2], &out_bytes[2 * sizeof(__m256i)]);
2320  storeu_avx2(h_vecs[3], &out_bytes[3 * sizeof(__m256i)]);
2321  storeu_avx2(h_vecs[4], &out_bytes[4 * sizeof(__m256i)]);
2322  storeu_avx2(h_vecs[5], &out_bytes[5 * sizeof(__m256i)]);
2323  storeu_avx2(h_vecs[6], &out_bytes[6 * sizeof(__m256i)]);
2324  storeu_avx2(h_vecs[7], &out_bytes[7 * sizeof(__m256i)]);
2325}
2326
2327void
2328blake3_hash_many_avx2(
2329    const uint8_t *const *inputs,
2330    size_t                num_inputs,
2331    size_t                blocks,
2332    const uint32_t        key[8],
2333    uint64_t              counter,
2334    int                   increment_counter,
2335    uint8_t               flags,
2336    uint8_t               flags_start,
2337    uint8_t               flags_end,
2338    uint8_t              *out_bytes
2339)
2340{
2341  while (num_inputs >= DEGREE_AVX2) {
2342    blake3_hash8_avx2(
2343        inputs, blocks, key, counter, increment_counter, flags, flags_start, flags_end, out_bytes
2344    );
2345    if (increment_counter) {
2346      counter += DEGREE_AVX2;
2347    }
2348    inputs += DEGREE_AVX2;
2349    num_inputs -= DEGREE_AVX2;
2350    out_bytes = &out_bytes[DEGREE_AVX2 * BLAKE3_OUT_LEN];
2351  }
2352  blake3_hash_many_sse41(
2353      inputs,
2354      num_inputs,
2355      blocks,
2356      key,
2357      counter,
2358      increment_counter,
2359      flags,
2360      flags_start,
2361      flags_end,
2362      out_bytes
2363  );
2364}
2365#if defined(__clang__)
2366#pragma clang attribute pop
2367#elif defined(__GNUC__)
2368#pragma GCC pop_options
2369#endif
2370
2371#if defined(__clang__)
2372#pragma clang attribute push(__attribute__((target("avx512f,avx512vl"))), apply_to = function)
2373#elif defined(__GNUC__)
2374#pragma GCC push_options
2375#pragma GCC target("avx512f,avx512vl")
2376#endif
2377static inline __m128i
2378loadu_128_avx512(const uint8_t src[16])
2379{
2380  return _mm_loadu_si128((const __m128i *)src);
2381}
2382
2383static inline __m256i
2384loadu_256_avx512(const uint8_t src[32])
2385{
2386  return _mm256_loadu_si256((const __m256i *)src);
2387}
2388
2389static inline __m512i
2390loadu_512_avx512(const uint8_t src[64])
2391{
2392  return _mm512_loadu_si512((const __m512i *)src);
2393}
2394
2395static inline void
2396storeu_128_avx512(__m128i src, uint8_t dest[16])
2397{
2398  _mm_storeu_si128((__m128i *)dest, src);
2399}
2400
2401static inline void
2402storeu_256_avx512(__m256i src, uint8_t dest[32])
2403{
2404  _mm256_storeu_si256((__m256i *)dest, src);
2405}
2406
2407static inline __m128i
2408add_128_avx512(__m128i a, __m128i b)
2409{
2410  return _mm_add_epi32(a, b);
2411}
2412
2413static inline __m256i
2414add_256_avx512(__m256i a, __m256i b)
2415{
2416  return _mm256_add_epi32(a, b);
2417}
2418
2419static inline __m512i
2420add_512_avx512(__m512i a, __m512i b)
2421{
2422  return _mm512_add_epi32(a, b);
2423}
2424
2425static inline __m128i
2426xor_128_avx512(__m128i a, __m128i b)
2427{
2428  return _mm_xor_si128(a, b);
2429}
2430
2431static inline __m256i
2432xor_256_avx512(__m256i a, __m256i b)
2433{
2434  return _mm256_xor_si256(a, b);
2435}
2436
2437static inline __m512i
2438xor_512_avx512(__m512i a, __m512i b)
2439{
2440  return _mm512_xor_si512(a, b);
2441}
2442
2443static inline __m128i
2444set1_128_avx512(uint32_t x)
2445{
2446  return _mm_set1_epi32((int32_t)x);
2447}
2448
2449static inline __m256i
2450set1_256_avx512(uint32_t x)
2451{
2452  return _mm256_set1_epi32((int32_t)x);
2453}
2454
2455static inline __m512i
2456set1_512_avx512(uint32_t x)
2457{
2458  return _mm512_set1_epi32((int32_t)x);
2459}
2460
2461static inline __m128i
2462set4_avx512(uint32_t a, uint32_t b, uint32_t c, uint32_t d)
2463{
2464  return _mm_setr_epi32((int32_t)a, (int32_t)b, (int32_t)c, (int32_t)d);
2465}
2466
2467static inline __m128i
2468rot16_128_avx512(__m128i x)
2469{
2470  return _mm_ror_epi32(x, 16);
2471}
2472
2473static inline __m256i
2474rot16_256_avx512(__m256i x)
2475{
2476  return _mm256_ror_epi32(x, 16);
2477}
2478
2479static inline __m512i
2480rot16_512_avx512(__m512i x)
2481{
2482  return _mm512_ror_epi32(x, 16);
2483}
2484
2485static inline __m128i
2486rot12_128_avx512(__m128i x)
2487{
2488  return _mm_ror_epi32(x, 12);
2489}
2490
2491static inline __m256i
2492rot12_256_avx512(__m256i x)
2493{
2494  return _mm256_ror_epi32(x, 12);
2495}
2496
2497static inline __m512i
2498rot12_512_avx512(__m512i x)
2499{
2500  return _mm512_ror_epi32(x, 12);
2501}
2502
2503static inline __m128i
2504rot8_128_avx512(__m128i x)
2505{
2506  return _mm_ror_epi32(x, 8);
2507}
2508
2509static inline __m256i
2510rot8_256_avx512(__m256i x)
2511{
2512  return _mm256_ror_epi32(x, 8);
2513}
2514
2515static inline __m512i
2516rot8_512_avx512(__m512i x)
2517{
2518  return _mm512_ror_epi32(x, 8);
2519}
2520
2521static inline __m128i
2522rot7_128_avx512(__m128i x)
2523{
2524  return _mm_ror_epi32(x, 7);
2525}
2526
2527static inline __m256i
2528rot7_256_avx512(__m256i x)
2529{
2530  return _mm256_ror_epi32(x, 7);
2531}
2532
2533static inline __m512i
2534rot7_512_avx512(__m512i x)
2535{
2536  return _mm512_ror_epi32(x, 7);
2537}
2538
2539static inline void
2540g1_avx512(__m128i *row0, __m128i *row1, __m128i *row2, __m128i *row3, __m128i m)
2541{
2542  *row0 = add_128_avx512(add_128_avx512(*row0, m), *row1);
2543  *row3 = xor_128_avx512(*row3, *row0);
2544  *row3 = rot16_128_avx512(*row3);
2545  *row2 = add_128_avx512(*row2, *row3);
2546  *row1 = xor_128_avx512(*row1, *row2);
2547  *row1 = rot12_128_avx512(*row1);
2548}
2549
2550static inline void
2551g2_avx512(__m128i *row0, __m128i *row1, __m128i *row2, __m128i *row3, __m128i m)
2552{
2553  *row0 = add_128_avx512(add_128_avx512(*row0, m), *row1);
2554  *row3 = xor_128_avx512(*row3, *row0);
2555  *row3 = rot8_128_avx512(*row3);
2556  *row2 = add_128_avx512(*row2, *row3);
2557  *row1 = xor_128_avx512(*row1, *row2);
2558  *row1 = rot7_128_avx512(*row1);
2559}
2560
2561static inline void
2562diagonalize_avx512(__m128i *row0, __m128i *row2, __m128i *row3)
2563{
2564  *row0 = _mm_shuffle_epi32(*row0, _MM_SHUFFLE(2, 1, 0, 3));
2565  *row3 = _mm_shuffle_epi32(*row3, _MM_SHUFFLE(1, 0, 3, 2));
2566  *row2 = _mm_shuffle_epi32(*row2, _MM_SHUFFLE(0, 3, 2, 1));
2567}
2568
2569static inline void
2570undiagonalize_avx512(__m128i *row0, __m128i *row2, __m128i *row3)
2571{
2572  *row0 = _mm_shuffle_epi32(*row0, _MM_SHUFFLE(0, 3, 2, 1));
2573  *row3 = _mm_shuffle_epi32(*row3, _MM_SHUFFLE(1, 0, 3, 2));
2574  *row2 = _mm_shuffle_epi32(*row2, _MM_SHUFFLE(2, 1, 0, 3));
2575}
2576
2577static inline void
2578compress_pre_avx512(
2579    __m128i        rows[4],
2580    const uint32_t cv[8],
2581    const uint8_t  block[BLAKE3_BLOCK_LEN],
2582    uint8_t        block_len,
2583    uint64_t       counter,
2584    uint8_t        flags
2585)
2586{
2587  __m128i m0, m1, m2, m3;
2588  __m128i t0, t1, t2, t3;
2589
2590  rows[0] = loadu_128_avx512((const uint8_t *)&cv[0]);
2591  rows[1] = loadu_128_avx512((const uint8_t *)&cv[4]);
2592  rows[2] = set4_avx512(IV[0], IV[1], IV[2], IV[3]);
2593  rows[3] = set4_avx512(
2594      counter_low(counter), counter_high(counter), (uint32_t)block_len, (uint32_t)flags
2595  );
2596
2597  m0 = loadu_128_avx512(&block[sizeof(__m128i) * 0]);
2598  m1 = loadu_128_avx512(&block[sizeof(__m128i) * 1]);
2599  m2 = loadu_128_avx512(&block[sizeof(__m128i) * 2]);
2600  m3 = loadu_128_avx512(&block[sizeof(__m128i) * 3]);
2601
2602  /* round 1 */
2603  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(2, 0, 2, 0));
2604  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2605  t1 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 3, 1));
2606  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2607  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2608  t2 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(2, 0, 2, 0));
2609  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(2, 1, 0, 3));
2610  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2611  t3 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 1, 3, 1));
2612  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(2, 1, 0, 3));
2613  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2614  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2615  m0 = t0;
2616  m1 = t1;
2617  m2 = t2;
2618  m3 = t3;
2619
2620  /* round 2 */
2621  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
2622  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
2623  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2624  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
2625  t1 = _mm_blend_epi16(m0, t1, 0xCC);
2626  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2627  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2628  t2 = _mm_unpacklo_epi64(m3, m1);
2629  t2 = _mm_blend_epi16(t2, m2, 0xC0);
2630  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(1, 3, 2, 0));
2631  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2632  t3 = _mm_unpackhi_epi32(m1, m3);
2633  t3 = _mm_unpacklo_epi32(m2, t3);
2634  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(0, 1, 3, 2));
2635  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2636  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2637  m0 = t0;
2638  m1 = t1;
2639  m2 = t2;
2640  m3 = t3;
2641
2642  /* round 3 */
2643  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
2644  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
2645  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2646  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
2647  t1 = _mm_blend_epi16(m0, t1, 0xCC);
2648  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2649  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2650  t2 = _mm_unpacklo_epi64(m3, m1);
2651  t2 = _mm_blend_epi16(t2, m2, 0xC0);
2652  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(1, 3, 2, 0));
2653  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2654  t3 = _mm_unpackhi_epi32(m1, m3);
2655  t3 = _mm_unpacklo_epi32(m2, t3);
2656  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(0, 1, 3, 2));
2657  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2658  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2659  m0 = t0;
2660  m1 = t1;
2661  m2 = t2;
2662  m3 = t3;
2663
2664  /* round 4 */
2665  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
2666  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
2667  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2668  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
2669  t1 = _mm_blend_epi16(m0, t1, 0xCC);
2670  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2671  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2672  t2 = _mm_unpacklo_epi64(m3, m1);
2673  t2 = _mm_blend_epi16(t2, m2, 0xC0);
2674  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(1, 3, 2, 0));
2675  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2676  t3 = _mm_unpackhi_epi32(m1, m3);
2677  t3 = _mm_unpacklo_epi32(m2, t3);
2678  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(0, 1, 3, 2));
2679  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2680  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2681  m0 = t0;
2682  m1 = t1;
2683  m2 = t2;
2684  m3 = t3;
2685
2686  /* round 5 */
2687  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
2688  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
2689  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2690  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
2691  t1 = _mm_blend_epi16(m0, t1, 0xCC);
2692  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2693  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2694  t2 = _mm_unpacklo_epi64(m3, m1);
2695  t2 = _mm_blend_epi16(t2, m2, 0xC0);
2696  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(1, 3, 2, 0));
2697  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2698  t3 = _mm_unpackhi_epi32(m1, m3);
2699  t3 = _mm_unpacklo_epi32(m2, t3);
2700  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(0, 1, 3, 2));
2701  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2702  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2703  m0 = t0;
2704  m1 = t1;
2705  m2 = t2;
2706  m3 = t3;
2707
2708  /* round 6 */
2709  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
2710  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
2711  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2712  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
2713  t1 = _mm_blend_epi16(m0, t1, 0xCC);
2714  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2715  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2716  t2 = _mm_unpacklo_epi64(m3, m1);
2717  t2 = _mm_blend_epi16(t2, m2, 0xC0);
2718  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(1, 3, 2, 0));
2719  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2720  t3 = _mm_unpackhi_epi32(m1, m3);
2721  t3 = _mm_unpacklo_epi32(m2, t3);
2722  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(0, 1, 3, 2));
2723  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2724  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2725  m0 = t0;
2726  m1 = t1;
2727  m2 = t2;
2728  m3 = t3;
2729
2730  /* round 7 */
2731  t0 = _mm_shuffle_ps2(m0, m1, _MM_SHUFFLE(3, 1, 1, 2));
2732  t0 = _mm_shuffle_epi32(t0, _MM_SHUFFLE(0, 3, 2, 1));
2733  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t0);
2734  t1 = _mm_shuffle_ps2(m2, m3, _MM_SHUFFLE(3, 3, 2, 2));
2735  t1 = _mm_blend_epi16(m0, t1, 0xCC);
2736  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t1);
2737  diagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2738  t2 = _mm_unpacklo_epi64(m3, m1);
2739  t2 = _mm_blend_epi16(t2, m2, 0xC0);
2740  t2 = _mm_shuffle_epi32(t2, _MM_SHUFFLE(1, 3, 2, 0));
2741  g1_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t2);
2742  t3 = _mm_unpackhi_epi32(m1, m3);
2743  t3 = _mm_unpacklo_epi32(m2, t3);
2744  t3 = _mm_shuffle_epi32(t3, _MM_SHUFFLE(0, 1, 3, 2));
2745  g2_avx512(&rows[0], &rows[1], &rows[2], &rows[3], t3);
2746  undiagonalize_avx512(&rows[0], &rows[2], &rows[3]);
2747}
2748
2749void
2750blake3_compress_in_place_avx512(
2751    uint32_t      cv[8],
2752    const uint8_t block[BLAKE3_BLOCK_LEN],
2753    uint8_t       block_len,
2754    uint64_t      counter,
2755    uint8_t       flags
2756)
2757{
2758  __m128i rows[4];
2759
2760  compress_pre_avx512(rows, cv, block, block_len, counter, flags);
2761  storeu_128_avx512(xor_128_avx512(rows[0], rows[2]), (uint8_t *)&cv[0]);
2762  storeu_128_avx512(xor_128_avx512(rows[1], rows[3]), (uint8_t *)&cv[4]);
2763}
2764
2765void
2766blake3_compress_xof_avx512(
2767    const uint32_t cv[8],
2768    const uint8_t  block[BLAKE3_BLOCK_LEN],
2769    uint8_t        block_len,
2770    uint64_t       counter,
2771    uint8_t        flags,
2772    uint8_t        out_buf[64]
2773)
2774{
2775  __m128i rows[4];
2776
2777  compress_pre_avx512(rows, cv, block, block_len, counter, flags);
2778  storeu_128_avx512(xor_128_avx512(rows[0], rows[2]), &out_buf[0]);
2779  storeu_128_avx512(xor_128_avx512(rows[1], rows[3]), &out_buf[16]);
2780  storeu_128_avx512(
2781      xor_128_avx512(rows[2], loadu_128_avx512((const uint8_t *)&cv[0])), &out_buf[32]
2782  );
2783  storeu_128_avx512(
2784      xor_128_avx512(rows[3], loadu_128_avx512((const uint8_t *)&cv[4])), &out_buf[48]
2785  );
2786}
2787
2788static inline void
2789round_fn4_avx512(__m128i v[16], __m128i m[16], size_t r)
2790{
2791  v[0]  = add_128_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][0]]);
2792  v[1]  = add_128_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][2]]);
2793  v[2]  = add_128_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][4]]);
2794  v[3]  = add_128_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][6]]);
2795  v[0]  = add_128_avx512(v[0], v[4]);
2796  v[1]  = add_128_avx512(v[1], v[5]);
2797  v[2]  = add_128_avx512(v[2], v[6]);
2798  v[3]  = add_128_avx512(v[3], v[7]);
2799  v[12] = xor_128_avx512(v[12], v[0]);
2800  v[13] = xor_128_avx512(v[13], v[1]);
2801  v[14] = xor_128_avx512(v[14], v[2]);
2802  v[15] = xor_128_avx512(v[15], v[3]);
2803  v[12] = rot16_128_avx512(v[12]);
2804  v[13] = rot16_128_avx512(v[13]);
2805  v[14] = rot16_128_avx512(v[14]);
2806  v[15] = rot16_128_avx512(v[15]);
2807  v[8]  = add_128_avx512(v[8], v[12]);
2808  v[9]  = add_128_avx512(v[9], v[13]);
2809  v[10] = add_128_avx512(v[10], v[14]);
2810  v[11] = add_128_avx512(v[11], v[15]);
2811  v[4]  = xor_128_avx512(v[4], v[8]);
2812  v[5]  = xor_128_avx512(v[5], v[9]);
2813  v[6]  = xor_128_avx512(v[6], v[10]);
2814  v[7]  = xor_128_avx512(v[7], v[11]);
2815  v[4]  = rot12_128_avx512(v[4]);
2816  v[5]  = rot12_128_avx512(v[5]);
2817  v[6]  = rot12_128_avx512(v[6]);
2818  v[7]  = rot12_128_avx512(v[7]);
2819  v[0]  = add_128_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][1]]);
2820  v[1]  = add_128_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][3]]);
2821  v[2]  = add_128_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][5]]);
2822  v[3]  = add_128_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][7]]);
2823  v[0]  = add_128_avx512(v[0], v[4]);
2824  v[1]  = add_128_avx512(v[1], v[5]);
2825  v[2]  = add_128_avx512(v[2], v[6]);
2826  v[3]  = add_128_avx512(v[3], v[7]);
2827  v[12] = xor_128_avx512(v[12], v[0]);
2828  v[13] = xor_128_avx512(v[13], v[1]);
2829  v[14] = xor_128_avx512(v[14], v[2]);
2830  v[15] = xor_128_avx512(v[15], v[3]);
2831  v[12] = rot8_128_avx512(v[12]);
2832  v[13] = rot8_128_avx512(v[13]);
2833  v[14] = rot8_128_avx512(v[14]);
2834  v[15] = rot8_128_avx512(v[15]);
2835  v[8]  = add_128_avx512(v[8], v[12]);
2836  v[9]  = add_128_avx512(v[9], v[13]);
2837  v[10] = add_128_avx512(v[10], v[14]);
2838  v[11] = add_128_avx512(v[11], v[15]);
2839  v[4]  = xor_128_avx512(v[4], v[8]);
2840  v[5]  = xor_128_avx512(v[5], v[9]);
2841  v[6]  = xor_128_avx512(v[6], v[10]);
2842  v[7]  = xor_128_avx512(v[7], v[11]);
2843  v[4]  = rot7_128_avx512(v[4]);
2844  v[5]  = rot7_128_avx512(v[5]);
2845  v[6]  = rot7_128_avx512(v[6]);
2846  v[7]  = rot7_128_avx512(v[7]);
2847
2848  v[0]  = add_128_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][8]]);
2849  v[1]  = add_128_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][10]]);
2850  v[2]  = add_128_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][12]]);
2851  v[3]  = add_128_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][14]]);
2852  v[0]  = add_128_avx512(v[0], v[5]);
2853  v[1]  = add_128_avx512(v[1], v[6]);
2854  v[2]  = add_128_avx512(v[2], v[7]);
2855  v[3]  = add_128_avx512(v[3], v[4]);
2856  v[15] = xor_128_avx512(v[15], v[0]);
2857  v[12] = xor_128_avx512(v[12], v[1]);
2858  v[13] = xor_128_avx512(v[13], v[2]);
2859  v[14] = xor_128_avx512(v[14], v[3]);
2860  v[15] = rot16_128_avx512(v[15]);
2861  v[12] = rot16_128_avx512(v[12]);
2862  v[13] = rot16_128_avx512(v[13]);
2863  v[14] = rot16_128_avx512(v[14]);
2864  v[10] = add_128_avx512(v[10], v[15]);
2865  v[11] = add_128_avx512(v[11], v[12]);
2866  v[8]  = add_128_avx512(v[8], v[13]);
2867  v[9]  = add_128_avx512(v[9], v[14]);
2868  v[5]  = xor_128_avx512(v[5], v[10]);
2869  v[6]  = xor_128_avx512(v[6], v[11]);
2870  v[7]  = xor_128_avx512(v[7], v[8]);
2871  v[4]  = xor_128_avx512(v[4], v[9]);
2872  v[5]  = rot12_128_avx512(v[5]);
2873  v[6]  = rot12_128_avx512(v[6]);
2874  v[7]  = rot12_128_avx512(v[7]);
2875  v[4]  = rot12_128_avx512(v[4]);
2876  v[0]  = add_128_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][9]]);
2877  v[1]  = add_128_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][11]]);
2878  v[2]  = add_128_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][13]]);
2879  v[3]  = add_128_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][15]]);
2880  v[0]  = add_128_avx512(v[0], v[5]);
2881  v[1]  = add_128_avx512(v[1], v[6]);
2882  v[2]  = add_128_avx512(v[2], v[7]);
2883  v[3]  = add_128_avx512(v[3], v[4]);
2884  v[15] = xor_128_avx512(v[15], v[0]);
2885  v[12] = xor_128_avx512(v[12], v[1]);
2886  v[13] = xor_128_avx512(v[13], v[2]);
2887  v[14] = xor_128_avx512(v[14], v[3]);
2888  v[15] = rot8_128_avx512(v[15]);
2889  v[12] = rot8_128_avx512(v[12]);
2890  v[13] = rot8_128_avx512(v[13]);
2891  v[14] = rot8_128_avx512(v[14]);
2892  v[10] = add_128_avx512(v[10], v[15]);
2893  v[11] = add_128_avx512(v[11], v[12]);
2894  v[8]  = add_128_avx512(v[8], v[13]);
2895  v[9]  = add_128_avx512(v[9], v[14]);
2896  v[5]  = xor_128_avx512(v[5], v[10]);
2897  v[6]  = xor_128_avx512(v[6], v[11]);
2898  v[7]  = xor_128_avx512(v[7], v[8]);
2899  v[4]  = xor_128_avx512(v[4], v[9]);
2900  v[5]  = rot7_128_avx512(v[5]);
2901  v[6]  = rot7_128_avx512(v[6]);
2902  v[7]  = rot7_128_avx512(v[7]);
2903  v[4]  = rot7_128_avx512(v[4]);
2904}
2905
2906static inline void
2907round_fn8_avx512(__m256i v[16], __m256i m[16], size_t r)
2908{
2909  v[0]  = add_256_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][0]]);
2910  v[1]  = add_256_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][2]]);
2911  v[2]  = add_256_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][4]]);
2912  v[3]  = add_256_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][6]]);
2913  v[0]  = add_256_avx512(v[0], v[4]);
2914  v[1]  = add_256_avx512(v[1], v[5]);
2915  v[2]  = add_256_avx512(v[2], v[6]);
2916  v[3]  = add_256_avx512(v[3], v[7]);
2917  v[12] = xor_256_avx512(v[12], v[0]);
2918  v[13] = xor_256_avx512(v[13], v[1]);
2919  v[14] = xor_256_avx512(v[14], v[2]);
2920  v[15] = xor_256_avx512(v[15], v[3]);
2921  v[12] = rot16_256_avx512(v[12]);
2922  v[13] = rot16_256_avx512(v[13]);
2923  v[14] = rot16_256_avx512(v[14]);
2924  v[15] = rot16_256_avx512(v[15]);
2925  v[8]  = add_256_avx512(v[8], v[12]);
2926  v[9]  = add_256_avx512(v[9], v[13]);
2927  v[10] = add_256_avx512(v[10], v[14]);
2928  v[11] = add_256_avx512(v[11], v[15]);
2929  v[4]  = xor_256_avx512(v[4], v[8]);
2930  v[5]  = xor_256_avx512(v[5], v[9]);
2931  v[6]  = xor_256_avx512(v[6], v[10]);
2932  v[7]  = xor_256_avx512(v[7], v[11]);
2933  v[4]  = rot12_256_avx512(v[4]);
2934  v[5]  = rot12_256_avx512(v[5]);
2935  v[6]  = rot12_256_avx512(v[6]);
2936  v[7]  = rot12_256_avx512(v[7]);
2937  v[0]  = add_256_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][1]]);
2938  v[1]  = add_256_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][3]]);
2939  v[2]  = add_256_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][5]]);
2940  v[3]  = add_256_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][7]]);
2941  v[0]  = add_256_avx512(v[0], v[4]);
2942  v[1]  = add_256_avx512(v[1], v[5]);
2943  v[2]  = add_256_avx512(v[2], v[6]);
2944  v[3]  = add_256_avx512(v[3], v[7]);
2945  v[12] = xor_256_avx512(v[12], v[0]);
2946  v[13] = xor_256_avx512(v[13], v[1]);
2947  v[14] = xor_256_avx512(v[14], v[2]);
2948  v[15] = xor_256_avx512(v[15], v[3]);
2949  v[12] = rot8_256_avx512(v[12]);
2950  v[13] = rot8_256_avx512(v[13]);
2951  v[14] = rot8_256_avx512(v[14]);
2952  v[15] = rot8_256_avx512(v[15]);
2953  v[8]  = add_256_avx512(v[8], v[12]);
2954  v[9]  = add_256_avx512(v[9], v[13]);
2955  v[10] = add_256_avx512(v[10], v[14]);
2956  v[11] = add_256_avx512(v[11], v[15]);
2957  v[4]  = xor_256_avx512(v[4], v[8]);
2958  v[5]  = xor_256_avx512(v[5], v[9]);
2959  v[6]  = xor_256_avx512(v[6], v[10]);
2960  v[7]  = xor_256_avx512(v[7], v[11]);
2961  v[4]  = rot7_256_avx512(v[4]);
2962  v[5]  = rot7_256_avx512(v[5]);
2963  v[6]  = rot7_256_avx512(v[6]);
2964  v[7]  = rot7_256_avx512(v[7]);
2965
2966  v[0]  = add_256_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][8]]);
2967  v[1]  = add_256_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][10]]);
2968  v[2]  = add_256_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][12]]);
2969  v[3]  = add_256_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][14]]);
2970  v[0]  = add_256_avx512(v[0], v[5]);
2971  v[1]  = add_256_avx512(v[1], v[6]);
2972  v[2]  = add_256_avx512(v[2], v[7]);
2973  v[3]  = add_256_avx512(v[3], v[4]);
2974  v[15] = xor_256_avx512(v[15], v[0]);
2975  v[12] = xor_256_avx512(v[12], v[1]);
2976  v[13] = xor_256_avx512(v[13], v[2]);
2977  v[14] = xor_256_avx512(v[14], v[3]);
2978  v[15] = rot16_256_avx512(v[15]);
2979  v[12] = rot16_256_avx512(v[12]);
2980  v[13] = rot16_256_avx512(v[13]);
2981  v[14] = rot16_256_avx512(v[14]);
2982  v[10] = add_256_avx512(v[10], v[15]);
2983  v[11] = add_256_avx512(v[11], v[12]);
2984  v[8]  = add_256_avx512(v[8], v[13]);
2985  v[9]  = add_256_avx512(v[9], v[14]);
2986  v[5]  = xor_256_avx512(v[5], v[10]);
2987  v[6]  = xor_256_avx512(v[6], v[11]);
2988  v[7]  = xor_256_avx512(v[7], v[8]);
2989  v[4]  = xor_256_avx512(v[4], v[9]);
2990  v[5]  = rot12_256_avx512(v[5]);
2991  v[6]  = rot12_256_avx512(v[6]);
2992  v[7]  = rot12_256_avx512(v[7]);
2993  v[4]  = rot12_256_avx512(v[4]);
2994  v[0]  = add_256_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][9]]);
2995  v[1]  = add_256_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][11]]);
2996  v[2]  = add_256_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][13]]);
2997  v[3]  = add_256_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][15]]);
2998  v[0]  = add_256_avx512(v[0], v[5]);
2999  v[1]  = add_256_avx512(v[1], v[6]);
3000  v[2]  = add_256_avx512(v[2], v[7]);
3001  v[3]  = add_256_avx512(v[3], v[4]);
3002  v[15] = xor_256_avx512(v[15], v[0]);
3003  v[12] = xor_256_avx512(v[12], v[1]);
3004  v[13] = xor_256_avx512(v[13], v[2]);
3005  v[14] = xor_256_avx512(v[14], v[3]);
3006  v[15] = rot8_256_avx512(v[15]);
3007  v[12] = rot8_256_avx512(v[12]);
3008  v[13] = rot8_256_avx512(v[13]);
3009  v[14] = rot8_256_avx512(v[14]);
3010  v[10] = add_256_avx512(v[10], v[15]);
3011  v[11] = add_256_avx512(v[11], v[12]);
3012  v[8]  = add_256_avx512(v[8], v[13]);
3013  v[9]  = add_256_avx512(v[9], v[14]);
3014  v[5]  = xor_256_avx512(v[5], v[10]);
3015  v[6]  = xor_256_avx512(v[6], v[11]);
3016  v[7]  = xor_256_avx512(v[7], v[8]);
3017  v[4]  = xor_256_avx512(v[4], v[9]);
3018  v[5]  = rot7_256_avx512(v[5]);
3019  v[6]  = rot7_256_avx512(v[6]);
3020  v[7]  = rot7_256_avx512(v[7]);
3021  v[4]  = rot7_256_avx512(v[4]);
3022}
3023
3024static inline void
3025round_fn16_avx512(__m512i v[16], __m512i m[16], size_t r)
3026{
3027  v[0]  = add_512_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][0]]);
3028  v[1]  = add_512_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][2]]);
3029  v[2]  = add_512_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][4]]);
3030  v[3]  = add_512_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][6]]);
3031  v[0]  = add_512_avx512(v[0], v[4]);
3032  v[1]  = add_512_avx512(v[1], v[5]);
3033  v[2]  = add_512_avx512(v[2], v[6]);
3034  v[3]  = add_512_avx512(v[3], v[7]);
3035  v[12] = xor_512_avx512(v[12], v[0]);
3036  v[13] = xor_512_avx512(v[13], v[1]);
3037  v[14] = xor_512_avx512(v[14], v[2]);
3038  v[15] = xor_512_avx512(v[15], v[3]);
3039  v[12] = rot16_512_avx512(v[12]);
3040  v[13] = rot16_512_avx512(v[13]);
3041  v[14] = rot16_512_avx512(v[14]);
3042  v[15] = rot16_512_avx512(v[15]);
3043  v[8]  = add_512_avx512(v[8], v[12]);
3044  v[9]  = add_512_avx512(v[9], v[13]);
3045  v[10] = add_512_avx512(v[10], v[14]);
3046  v[11] = add_512_avx512(v[11], v[15]);
3047  v[4]  = xor_512_avx512(v[4], v[8]);
3048  v[5]  = xor_512_avx512(v[5], v[9]);
3049  v[6]  = xor_512_avx512(v[6], v[10]);
3050  v[7]  = xor_512_avx512(v[7], v[11]);
3051  v[4]  = rot12_512_avx512(v[4]);
3052  v[5]  = rot12_512_avx512(v[5]);
3053  v[6]  = rot12_512_avx512(v[6]);
3054  v[7]  = rot12_512_avx512(v[7]);
3055  v[0]  = add_512_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][1]]);
3056  v[1]  = add_512_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][3]]);
3057  v[2]  = add_512_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][5]]);
3058  v[3]  = add_512_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][7]]);
3059  v[0]  = add_512_avx512(v[0], v[4]);
3060  v[1]  = add_512_avx512(v[1], v[5]);
3061  v[2]  = add_512_avx512(v[2], v[6]);
3062  v[3]  = add_512_avx512(v[3], v[7]);
3063  v[12] = xor_512_avx512(v[12], v[0]);
3064  v[13] = xor_512_avx512(v[13], v[1]);
3065  v[14] = xor_512_avx512(v[14], v[2]);
3066  v[15] = xor_512_avx512(v[15], v[3]);
3067  v[12] = rot8_512_avx512(v[12]);
3068  v[13] = rot8_512_avx512(v[13]);
3069  v[14] = rot8_512_avx512(v[14]);
3070  v[15] = rot8_512_avx512(v[15]);
3071  v[8]  = add_512_avx512(v[8], v[12]);
3072  v[9]  = add_512_avx512(v[9], v[13]);
3073  v[10] = add_512_avx512(v[10], v[14]);
3074  v[11] = add_512_avx512(v[11], v[15]);
3075  v[4]  = xor_512_avx512(v[4], v[8]);
3076  v[5]  = xor_512_avx512(v[5], v[9]);
3077  v[6]  = xor_512_avx512(v[6], v[10]);
3078  v[7]  = xor_512_avx512(v[7], v[11]);
3079  v[4]  = rot7_512_avx512(v[4]);
3080  v[5]  = rot7_512_avx512(v[5]);
3081  v[6]  = rot7_512_avx512(v[6]);
3082  v[7]  = rot7_512_avx512(v[7]);
3083
3084  v[0]  = add_512_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][8]]);
3085  v[1]  = add_512_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][10]]);
3086  v[2]  = add_512_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][12]]);
3087  v[3]  = add_512_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][14]]);
3088  v[0]  = add_512_avx512(v[0], v[5]);
3089  v[1]  = add_512_avx512(v[1], v[6]);
3090  v[2]  = add_512_avx512(v[2], v[7]);
3091  v[3]  = add_512_avx512(v[3], v[4]);
3092  v[15] = xor_512_avx512(v[15], v[0]);
3093  v[12] = xor_512_avx512(v[12], v[1]);
3094  v[13] = xor_512_avx512(v[13], v[2]);
3095  v[14] = xor_512_avx512(v[14], v[3]);
3096  v[15] = rot16_512_avx512(v[15]);
3097  v[12] = rot16_512_avx512(v[12]);
3098  v[13] = rot16_512_avx512(v[13]);
3099  v[14] = rot16_512_avx512(v[14]);
3100  v[10] = add_512_avx512(v[10], v[15]);
3101  v[11] = add_512_avx512(v[11], v[12]);
3102  v[8]  = add_512_avx512(v[8], v[13]);
3103  v[9]  = add_512_avx512(v[9], v[14]);
3104  v[5]  = xor_512_avx512(v[5], v[10]);
3105  v[6]  = xor_512_avx512(v[6], v[11]);
3106  v[7]  = xor_512_avx512(v[7], v[8]);
3107  v[4]  = xor_512_avx512(v[4], v[9]);
3108  v[5]  = rot12_512_avx512(v[5]);
3109  v[6]  = rot12_512_avx512(v[6]);
3110  v[7]  = rot12_512_avx512(v[7]);
3111  v[4]  = rot12_512_avx512(v[4]);
3112  v[0]  = add_512_avx512(v[0], m[(size_t)MSG_SCHEDULE[r][9]]);
3113  v[1]  = add_512_avx512(v[1], m[(size_t)MSG_SCHEDULE[r][11]]);
3114  v[2]  = add_512_avx512(v[2], m[(size_t)MSG_SCHEDULE[r][13]]);
3115  v[3]  = add_512_avx512(v[3], m[(size_t)MSG_SCHEDULE[r][15]]);
3116  v[0]  = add_512_avx512(v[0], v[5]);
3117  v[1]  = add_512_avx512(v[1], v[6]);
3118  v[2]  = add_512_avx512(v[2], v[7]);
3119  v[3]  = add_512_avx512(v[3], v[4]);
3120  v[15] = xor_512_avx512(v[15], v[0]);
3121  v[12] = xor_512_avx512(v[12], v[1]);
3122  v[13] = xor_512_avx512(v[13], v[2]);
3123  v[14] = xor_512_avx512(v[14], v[3]);
3124  v[15] = rot8_512_avx512(v[15]);
3125  v[12] = rot8_512_avx512(v[12]);
3126  v[13] = rot8_512_avx512(v[13]);
3127  v[14] = rot8_512_avx512(v[14]);
3128  v[10] = add_512_avx512(v[10], v[15]);
3129  v[11] = add_512_avx512(v[11], v[12]);
3130  v[8]  = add_512_avx512(v[8], v[13]);
3131  v[9]  = add_512_avx512(v[9], v[14]);
3132  v[5]  = xor_512_avx512(v[5], v[10]);
3133  v[6]  = xor_512_avx512(v[6], v[11]);
3134  v[7]  = xor_512_avx512(v[7], v[8]);
3135  v[4]  = xor_512_avx512(v[4], v[9]);
3136  v[5]  = rot7_512_avx512(v[5]);
3137  v[6]  = rot7_512_avx512(v[6]);
3138  v[7]  = rot7_512_avx512(v[7]);
3139  v[4]  = rot7_512_avx512(v[4]);
3140}
3141
3142#define LO_IMM8 0x88
3143#define HI_IMM8 0xdd
3144
3145static inline __m512i
3146unpack_lo_128_avx512(__m512i a, __m512i b)
3147{
3148  return _mm512_shuffle_i32x4(a, b, LO_IMM8);
3149}
3150
3151static inline __m512i
3152unpack_hi_128_avx512(__m512i a, __m512i b)
3153{
3154  return _mm512_shuffle_i32x4(a, b, HI_IMM8);
3155}
3156
3157static inline void
3158transpose_vecs_512_avx512(__m512i vecs[16])
3159{
3160  __m512i ab_0 = _mm512_unpacklo_epi32(vecs[0], vecs[1]);
3161  __m512i ab_2 = _mm512_unpackhi_epi32(vecs[0], vecs[1]);
3162  __m512i cd_0 = _mm512_unpacklo_epi32(vecs[2], vecs[3]);
3163  __m512i cd_2 = _mm512_unpackhi_epi32(vecs[2], vecs[3]);
3164  __m512i ef_0 = _mm512_unpacklo_epi32(vecs[4], vecs[5]);
3165  __m512i ef_2 = _mm512_unpackhi_epi32(vecs[4], vecs[5]);
3166  __m512i gh_0 = _mm512_unpacklo_epi32(vecs[6], vecs[7]);
3167  __m512i gh_2 = _mm512_unpackhi_epi32(vecs[6], vecs[7]);
3168  __m512i ij_0 = _mm512_unpacklo_epi32(vecs[8], vecs[9]);
3169  __m512i ij_2 = _mm512_unpackhi_epi32(vecs[8], vecs[9]);
3170  __m512i kl_0 = _mm512_unpacklo_epi32(vecs[10], vecs[11]);
3171  __m512i kl_2 = _mm512_unpackhi_epi32(vecs[10], vecs[11]);
3172  __m512i mn_0 = _mm512_unpacklo_epi32(vecs[12], vecs[13]);
3173  __m512i mn_2 = _mm512_unpackhi_epi32(vecs[12], vecs[13]);
3174  __m512i op_0 = _mm512_unpacklo_epi32(vecs[14], vecs[15]);
3175  __m512i op_2 = _mm512_unpackhi_epi32(vecs[14], vecs[15]);
3176
3177  __m512i abcd_0 = _mm512_unpacklo_epi64(ab_0, cd_0);
3178  __m512i abcd_1 = _mm512_unpackhi_epi64(ab_0, cd_0);
3179  __m512i abcd_2 = _mm512_unpacklo_epi64(ab_2, cd_2);
3180  __m512i abcd_3 = _mm512_unpackhi_epi64(ab_2, cd_2);
3181  __m512i efgh_0 = _mm512_unpacklo_epi64(ef_0, gh_0);
3182  __m512i efgh_1 = _mm512_unpackhi_epi64(ef_0, gh_0);
3183  __m512i efgh_2 = _mm512_unpacklo_epi64(ef_2, gh_2);
3184  __m512i efgh_3 = _mm512_unpackhi_epi64(ef_2, gh_2);
3185  __m512i ijkl_0 = _mm512_unpacklo_epi64(ij_0, kl_0);
3186  __m512i ijkl_1 = _mm512_unpackhi_epi64(ij_0, kl_0);
3187  __m512i ijkl_2 = _mm512_unpacklo_epi64(ij_2, kl_2);
3188  __m512i ijkl_3 = _mm512_unpackhi_epi64(ij_2, kl_2);
3189  __m512i mnop_0 = _mm512_unpacklo_epi64(mn_0, op_0);
3190  __m512i mnop_1 = _mm512_unpackhi_epi64(mn_0, op_0);
3191  __m512i mnop_2 = _mm512_unpacklo_epi64(mn_2, op_2);
3192  __m512i mnop_3 = _mm512_unpackhi_epi64(mn_2, op_2);
3193
3194  __m512i abcdefgh_0 = unpack_lo_128_avx512(abcd_0, efgh_0);
3195  __m512i abcdefgh_1 = unpack_lo_128_avx512(abcd_1, efgh_1);
3196  __m512i abcdefgh_2 = unpack_lo_128_avx512(abcd_2, efgh_2);
3197  __m512i abcdefgh_3 = unpack_lo_128_avx512(abcd_3, efgh_3);
3198  __m512i abcdefgh_4 = unpack_hi_128_avx512(abcd_0, efgh_0);
3199  __m512i abcdefgh_5 = unpack_hi_128_avx512(abcd_1, efgh_1);
3200  __m512i abcdefgh_6 = unpack_hi_128_avx512(abcd_2, efgh_2);
3201  __m512i abcdefgh_7 = unpack_hi_128_avx512(abcd_3, efgh_3);
3202  __m512i ijklmnop_0 = unpack_lo_128_avx512(ijkl_0, mnop_0);
3203  __m512i ijklmnop_1 = unpack_lo_128_avx512(ijkl_1, mnop_1);
3204  __m512i ijklmnop_2 = unpack_lo_128_avx512(ijkl_2, mnop_2);
3205  __m512i ijklmnop_3 = unpack_lo_128_avx512(ijkl_3, mnop_3);
3206  __m512i ijklmnop_4 = unpack_hi_128_avx512(ijkl_0, mnop_0);
3207  __m512i ijklmnop_5 = unpack_hi_128_avx512(ijkl_1, mnop_1);
3208  __m512i ijklmnop_6 = unpack_hi_128_avx512(ijkl_2, mnop_2);
3209  __m512i ijklmnop_7 = unpack_hi_128_avx512(ijkl_3, mnop_3);
3210
3211  vecs[0]  = unpack_lo_128_avx512(abcdefgh_0, ijklmnop_0);
3212  vecs[1]  = unpack_lo_128_avx512(abcdefgh_1, ijklmnop_1);
3213  vecs[2]  = unpack_lo_128_avx512(abcdefgh_2, ijklmnop_2);
3214  vecs[3]  = unpack_lo_128_avx512(abcdefgh_3, ijklmnop_3);
3215  vecs[4]  = unpack_lo_128_avx512(abcdefgh_4, ijklmnop_4);
3216  vecs[5]  = unpack_lo_128_avx512(abcdefgh_5, ijklmnop_5);
3217  vecs[6]  = unpack_lo_128_avx512(abcdefgh_6, ijklmnop_6);
3218  vecs[7]  = unpack_lo_128_avx512(abcdefgh_7, ijklmnop_7);
3219  vecs[8]  = unpack_hi_128_avx512(abcdefgh_0, ijklmnop_0);
3220  vecs[9]  = unpack_hi_128_avx512(abcdefgh_1, ijklmnop_1);
3221  vecs[10] = unpack_hi_128_avx512(abcdefgh_2, ijklmnop_2);
3222  vecs[11] = unpack_hi_128_avx512(abcdefgh_3, ijklmnop_3);
3223  vecs[12] = unpack_hi_128_avx512(abcdefgh_4, ijklmnop_4);
3224  vecs[13] = unpack_hi_128_avx512(abcdefgh_5, ijklmnop_5);
3225  vecs[14] = unpack_hi_128_avx512(abcdefgh_6, ijklmnop_6);
3226  vecs[15] = unpack_hi_128_avx512(abcdefgh_7, ijklmnop_7);
3227}
3228
3229static inline void
3230transpose_msg_vecs16_avx512(const uint8_t *const *inputs, size_t block_offset, __m512i out_msg[16])
3231{
3232  size_t i;
3233
3234  out_msg[0]  = loadu_512_avx512(&inputs[0][block_offset]);
3235  out_msg[1]  = loadu_512_avx512(&inputs[1][block_offset]);
3236  out_msg[2]  = loadu_512_avx512(&inputs[2][block_offset]);
3237  out_msg[3]  = loadu_512_avx512(&inputs[3][block_offset]);
3238  out_msg[4]  = loadu_512_avx512(&inputs[4][block_offset]);
3239  out_msg[5]  = loadu_512_avx512(&inputs[5][block_offset]);
3240  out_msg[6]  = loadu_512_avx512(&inputs[6][block_offset]);
3241  out_msg[7]  = loadu_512_avx512(&inputs[7][block_offset]);
3242  out_msg[8]  = loadu_512_avx512(&inputs[8][block_offset]);
3243  out_msg[9]  = loadu_512_avx512(&inputs[9][block_offset]);
3244  out_msg[10] = loadu_512_avx512(&inputs[10][block_offset]);
3245  out_msg[11] = loadu_512_avx512(&inputs[11][block_offset]);
3246  out_msg[12] = loadu_512_avx512(&inputs[12][block_offset]);
3247  out_msg[13] = loadu_512_avx512(&inputs[13][block_offset]);
3248  out_msg[14] = loadu_512_avx512(&inputs[14][block_offset]);
3249  out_msg[15] = loadu_512_avx512(&inputs[15][block_offset]);
3250
3251  for (i = 0; i < 16; i++) {
3252    _mm_prefetch((const void *)&inputs[i][block_offset + 256], _MM_HINT_T0);
3253  }
3254  transpose_vecs_512_avx512(out_msg);
3255}
3256
3257static inline void
3258load_counters16_avx512(uint64_t counter, int increment_counter, __m512i *out_lo, __m512i *out_hi)
3259{
3260  const __m512i mask   = _mm512_set1_epi32(-increment_counter);
3261  const __m512i deltas = _mm512_set_epi32(15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0);
3262  const __m512i masked_deltas = _mm512_and_si512(deltas, mask);
3263  const __m512i low_words = _mm512_add_epi32(_mm512_set1_epi32((int32_t)counter), masked_deltas);
3264  const __m512i carries =
3265      _mm512_srli_epi32(_mm512_andnot_si512(low_words, _mm512_set1_epi32((int32_t)counter)), 31);
3266  const __m512i high_words = _mm512_add_epi32(_mm512_set1_epi32((int32_t)(counter >> 32)), carries);
3267
3268  *out_lo = low_words;
3269  *out_hi = high_words;
3270}
3271
3272void
3273blake3_hash16_avx512(
3274    const uint8_t *const *inputs,
3275    size_t                blocks,
3276    const uint32_t        key[8],
3277    uint64_t              counter,
3278    int                   increment_counter,
3279    uint8_t               flags,
3280    uint8_t               flags_start,
3281    uint8_t               flags_end,
3282    uint8_t              *out_bytes
3283)
3284{
3285  __m512i h_vecs[8];
3286  __m512i counter_low_vec, counter_high_vec;
3287  uint8_t block_flags;
3288  size_t  block;
3289
3290  h_vecs[0] = set1_512_avx512(key[0]);
3291  h_vecs[1] = set1_512_avx512(key[1]);
3292  h_vecs[2] = set1_512_avx512(key[2]);
3293  h_vecs[3] = set1_512_avx512(key[3]);
3294  h_vecs[4] = set1_512_avx512(key[4]);
3295  h_vecs[5] = set1_512_avx512(key[5]);
3296  h_vecs[6] = set1_512_avx512(key[6]);
3297  h_vecs[7] = set1_512_avx512(key[7]);
3298
3299  load_counters16_avx512(counter, increment_counter, &counter_low_vec, &counter_high_vec);
3300  block_flags = flags | flags_start;
3301
3302  for (block = 0; block < blocks; block++) {
3303    __m512i block_len_vec;
3304    __m512i block_flags_vec;
3305    __m512i msg_vecs[16];
3306    __m512i v[16];
3307
3308    if (block + 1 == blocks) {
3309      block_flags |= flags_end;
3310    }
3311    block_len_vec   = set1_512_avx512(BLAKE3_BLOCK_LEN);
3312    block_flags_vec = set1_512_avx512(block_flags);
3313    transpose_msg_vecs16_avx512(inputs, block * BLAKE3_BLOCK_LEN, msg_vecs);
3314
3315    v[0]  = h_vecs[0];
3316    v[1]  = h_vecs[1];
3317    v[2]  = h_vecs[2];
3318    v[3]  = h_vecs[3];
3319    v[4]  = h_vecs[4];
3320    v[5]  = h_vecs[5];
3321    v[6]  = h_vecs[6];
3322    v[7]  = h_vecs[7];
3323    v[8]  = set1_512_avx512(IV[0]);
3324    v[9]  = set1_512_avx512(IV[1]);
3325    v[10] = set1_512_avx512(IV[2]);
3326    v[11] = set1_512_avx512(IV[3]);
3327    v[12] = counter_low_vec;
3328    v[13] = counter_high_vec;
3329    v[14] = block_len_vec;
3330    v[15] = block_flags_vec;
3331
3332    round_fn16_avx512(v, msg_vecs, 0);
3333    round_fn16_avx512(v, msg_vecs, 1);
3334    round_fn16_avx512(v, msg_vecs, 2);
3335    round_fn16_avx512(v, msg_vecs, 3);
3336    round_fn16_avx512(v, msg_vecs, 4);
3337    round_fn16_avx512(v, msg_vecs, 5);
3338    round_fn16_avx512(v, msg_vecs, 6);
3339
3340    h_vecs[0] = xor_512_avx512(v[0], v[8]);
3341    h_vecs[1] = xor_512_avx512(v[1], v[9]);
3342    h_vecs[2] = xor_512_avx512(v[2], v[10]);
3343    h_vecs[3] = xor_512_avx512(v[3], v[11]);
3344    h_vecs[4] = xor_512_avx512(v[4], v[12]);
3345    h_vecs[5] = xor_512_avx512(v[5], v[13]);
3346    h_vecs[6] = xor_512_avx512(v[6], v[14]);
3347    h_vecs[7] = xor_512_avx512(v[7], v[15]);
3348
3349    block_flags = flags;
3350  }
3351
3352  __m512i padded[16] = {
3353      h_vecs[0],
3354      h_vecs[1],
3355      h_vecs[2],
3356      h_vecs[3],
3357      h_vecs[4],
3358      h_vecs[5],
3359      h_vecs[6],
3360      h_vecs[7],
3361      set1_512_avx512(0),
3362      set1_512_avx512(0),
3363      set1_512_avx512(0),
3364      set1_512_avx512(0),
3365      set1_512_avx512(0),
3366      set1_512_avx512(0),
3367      set1_512_avx512(0),
3368      set1_512_avx512(0),
3369  };
3370  transpose_vecs_512_avx512(padded);
3371  _mm256_mask_storeu_epi32(
3372      &out_bytes[0 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[0])
3373  );
3374  _mm256_mask_storeu_epi32(
3375      &out_bytes[1 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[1])
3376  );
3377  _mm256_mask_storeu_epi32(
3378      &out_bytes[2 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[2])
3379  );
3380  _mm256_mask_storeu_epi32(
3381      &out_bytes[3 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[3])
3382  );
3383  _mm256_mask_storeu_epi32(
3384      &out_bytes[4 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[4])
3385  );
3386  _mm256_mask_storeu_epi32(
3387      &out_bytes[5 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[5])
3388  );
3389  _mm256_mask_storeu_epi32(
3390      &out_bytes[6 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[6])
3391  );
3392  _mm256_mask_storeu_epi32(
3393      &out_bytes[7 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[7])
3394  );
3395  _mm256_mask_storeu_epi32(
3396      &out_bytes[8 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[8])
3397  );
3398  _mm256_mask_storeu_epi32(
3399      &out_bytes[9 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[9])
3400  );
3401  _mm256_mask_storeu_epi32(
3402      &out_bytes[10 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[10])
3403  );
3404  _mm256_mask_storeu_epi32(
3405      &out_bytes[11 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[11])
3406  );
3407  _mm256_mask_storeu_epi32(
3408      &out_bytes[12 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[12])
3409  );
3410  _mm256_mask_storeu_epi32(
3411      &out_bytes[13 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[13])
3412  );
3413  _mm256_mask_storeu_epi32(
3414      &out_bytes[14 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[14])
3415  );
3416  _mm256_mask_storeu_epi32(
3417      &out_bytes[15 * sizeof(__m256i)], (__mmask8)-1, _mm512_castsi512_si256(padded[15])
3418  );
3419}
3420
3421static inline void
3422transpose_msg_vecs8_avx512(const uint8_t *const *inputs, size_t block_offset, __m256i out_msg[16])
3423{
3424  out_msg[0]  = loadu_256_avx512(&inputs[0][block_offset]);
3425  out_msg[1]  = loadu_256_avx512(&inputs[1][block_offset]);
3426  out_msg[2]  = loadu_256_avx512(&inputs[2][block_offset]);
3427  out_msg[3]  = loadu_256_avx512(&inputs[3][block_offset]);
3428  out_msg[4]  = loadu_256_avx512(&inputs[4][block_offset]);
3429  out_msg[5]  = loadu_256_avx512(&inputs[5][block_offset]);
3430  out_msg[6]  = loadu_256_avx512(&inputs[6][block_offset]);
3431  out_msg[7]  = loadu_256_avx512(&inputs[7][block_offset]);
3432  out_msg[8]  = loadu_256_avx512(&inputs[0][block_offset + 32]);
3433  out_msg[9]  = loadu_256_avx512(&inputs[1][block_offset + 32]);
3434  out_msg[10] = loadu_256_avx512(&inputs[2][block_offset + 32]);
3435  out_msg[11] = loadu_256_avx512(&inputs[3][block_offset + 32]);
3436  out_msg[12] = loadu_256_avx512(&inputs[4][block_offset + 32]);
3437  out_msg[13] = loadu_256_avx512(&inputs[5][block_offset + 32]);
3438  out_msg[14] = loadu_256_avx512(&inputs[6][block_offset + 32]);
3439  out_msg[15] = loadu_256_avx512(&inputs[7][block_offset + 32]);
3440
3441  transpose_vecs_avx2(out_msg);
3442}
3443
3444static inline void
3445load_counters8_avx512(uint64_t counter, int increment_counter, __m256i *out_lo, __m256i *out_hi)
3446{
3447  const __m256i mask = _mm256_set1_epi32(-increment_counter);
3448  const __m256i add0 = _mm256_set_epi32(7, 6, 5, 4, 3, 2, 1, 0);
3449  const __m256i add1 = _mm256_and_si256(mask, add0);
3450  __m256i       l    = _mm256_add_epi32(_mm256_set1_epi32((int32_t)counter), add1);
3451  __m256i       carry =
3452      _mm256_srli_epi32(_mm256_andnot_si256(l, _mm256_set1_epi32((int32_t)counter)), 31);
3453  __m256i h = _mm256_add_epi32(_mm256_set1_epi32((int32_t)(counter >> 32)), carry);
3454
3455  *out_lo = l;
3456  *out_hi = h;
3457}
3458
3459void
3460blake3_hash8_avx512(
3461    const uint8_t *const *inputs,
3462    size_t                blocks,
3463    const uint32_t        key[8],
3464    uint64_t              counter,
3465    int                   increment_counter,
3466    uint8_t               flags,
3467    uint8_t               flags_start,
3468    uint8_t               flags_end,
3469    uint8_t              *out_bytes
3470)
3471{
3472  __m256i h_vecs[8];
3473  __m256i counter_low_vec, counter_high_vec;
3474  uint8_t block_flags;
3475  size_t  block;
3476
3477  h_vecs[0] = set1_256_avx512(key[0]);
3478  h_vecs[1] = set1_256_avx512(key[1]);
3479  h_vecs[2] = set1_256_avx512(key[2]);
3480  h_vecs[3] = set1_256_avx512(key[3]);
3481  h_vecs[4] = set1_256_avx512(key[4]);
3482  h_vecs[5] = set1_256_avx512(key[5]);
3483  h_vecs[6] = set1_256_avx512(key[6]);
3484  h_vecs[7] = set1_256_avx512(key[7]);
3485
3486  load_counters8_avx512(counter, increment_counter, &counter_low_vec, &counter_high_vec);
3487  block_flags = flags | flags_start;
3488
3489  for (block = 0; block < blocks; block++) {
3490    __m256i block_len_vec;
3491    __m256i block_flags_vec;
3492    __m256i msg_vecs[16];
3493    __m256i v[16];
3494
3495    if (block + 1 == blocks) {
3496      block_flags |= flags_end;
3497    }
3498    block_len_vec   = set1_256_avx512(BLAKE3_BLOCK_LEN);
3499    block_flags_vec = set1_256_avx512(block_flags);
3500    transpose_msg_vecs8_avx512(inputs, block * BLAKE3_BLOCK_LEN, msg_vecs);
3501
3502    v[0]  = h_vecs[0];
3503    v[1]  = h_vecs[1];
3504    v[2]  = h_vecs[2];
3505    v[3]  = h_vecs[3];
3506    v[4]  = h_vecs[4];
3507    v[5]  = h_vecs[5];
3508    v[6]  = h_vecs[6];
3509    v[7]  = h_vecs[7];
3510    v[8]  = set1_256_avx512(IV[0]);
3511    v[9]  = set1_256_avx512(IV[1]);
3512    v[10] = set1_256_avx512(IV[2]);
3513    v[11] = set1_256_avx512(IV[3]);
3514    v[12] = counter_low_vec;
3515    v[13] = counter_high_vec;
3516    v[14] = block_len_vec;
3517    v[15] = block_flags_vec;
3518
3519    round_fn8_avx512(v, msg_vecs, 0);
3520    round_fn8_avx512(v, msg_vecs, 1);
3521    round_fn8_avx512(v, msg_vecs, 2);
3522    round_fn8_avx512(v, msg_vecs, 3);
3523    round_fn8_avx512(v, msg_vecs, 4);
3524    round_fn8_avx512(v, msg_vecs, 5);
3525    round_fn8_avx512(v, msg_vecs, 6);
3526
3527    h_vecs[0] = xor_256_avx512(v[0], v[8]);
3528    h_vecs[1] = xor_256_avx512(v[1], v[9]);
3529    h_vecs[2] = xor_256_avx512(v[2], v[10]);
3530    h_vecs[3] = xor_256_avx512(v[3], v[11]);
3531    h_vecs[4] = xor_256_avx512(v[4], v[12]);
3532    h_vecs[5] = xor_256_avx512(v[5], v[13]);
3533    h_vecs[6] = xor_256_avx512(v[6], v[14]);
3534    h_vecs[7] = xor_256_avx512(v[7], v[15]);
3535
3536    block_flags = flags;
3537  }
3538
3539  transpose_vecs_avx2(h_vecs);
3540  storeu_256_avx512(h_vecs[0], &out_bytes[0 * sizeof(__m256i)]);
3541  storeu_256_avx512(h_vecs[1], &out_bytes[1 * sizeof(__m256i)]);
3542  storeu_256_avx512(h_vecs[2], &out_bytes[2 * sizeof(__m256i)]);
3543  storeu_256_avx512(h_vecs[3], &out_bytes[3 * sizeof(__m256i)]);
3544  storeu_256_avx512(h_vecs[4], &out_bytes[4 * sizeof(__m256i)]);
3545  storeu_256_avx512(h_vecs[5], &out_bytes[5 * sizeof(__m256i)]);
3546  storeu_256_avx512(h_vecs[6], &out_bytes[6 * sizeof(__m256i)]);
3547  storeu_256_avx512(h_vecs[7], &out_bytes[7 * sizeof(__m256i)]);
3548}
3549
3550static inline void
3551transpose_msg_vecs4_avx512(const uint8_t *const *inputs, size_t block_offset, __m128i out_msg[16])
3552{
3553  out_msg[0]  = loadu_128_avx512(&inputs[0][block_offset]);
3554  out_msg[1]  = loadu_128_avx512(&inputs[1][block_offset]);
3555  out_msg[2]  = loadu_128_avx512(&inputs[2][block_offset]);
3556  out_msg[3]  = loadu_128_avx512(&inputs[3][block_offset]);
3557  out_msg[4]  = loadu_128_avx512(&inputs[0][block_offset + 16]);
3558  out_msg[5]  = loadu_128_avx512(&inputs[1][block_offset + 16]);
3559  out_msg[6]  = loadu_128_avx512(&inputs[2][block_offset + 16]);
3560  out_msg[7]  = loadu_128_avx512(&inputs[3][block_offset + 16]);
3561  out_msg[8]  = loadu_128_avx512(&inputs[0][block_offset + 32]);
3562  out_msg[9]  = loadu_128_avx512(&inputs[1][block_offset + 32]);
3563  out_msg[10] = loadu_128_avx512(&inputs[2][block_offset + 32]);
3564  out_msg[11] = loadu_128_avx512(&inputs[3][block_offset + 32]);
3565  out_msg[12] = loadu_128_avx512(&inputs[0][block_offset + 48]);
3566  out_msg[13] = loadu_128_avx512(&inputs[1][block_offset + 48]);
3567  out_msg[14] = loadu_128_avx512(&inputs[2][block_offset + 48]);
3568  out_msg[15] = loadu_128_avx512(&inputs[3][block_offset + 48]);
3569
3570  transpose_vecs_sse2(out_msg);
3571}
3572
3573static inline void
3574load_counters4_avx512(uint64_t counter, int increment_counter, __m128i *out_lo, __m128i *out_hi)
3575{
3576  const __m128i mask  = _mm_set1_epi32(-increment_counter);
3577  const __m128i add0  = _mm_set_epi32(3, 2, 1, 0);
3578  const __m128i add1  = _mm_and_si128(mask, add0);
3579  __m128i       l     = _mm_add_epi32(_mm_set1_epi32((int32_t)counter), add1);
3580  __m128i       carry = _mm_srli_epi32(_mm_andnot_si128(l, _mm_set1_epi32((int32_t)counter)), 31);
3581  __m128i       h     = _mm_add_epi32(_mm_set1_epi32((int32_t)(counter >> 32)), carry);
3582
3583  *out_lo = l;
3584  *out_hi = h;
3585}
3586
3587void
3588blake3_hash4_avx512(
3589    const uint8_t *const *inputs,
3590    size_t                blocks,
3591    const uint32_t        key[8],
3592    uint64_t              counter,
3593    int                   increment_counter,
3594    uint8_t               flags,
3595    uint8_t               flags_start,
3596    uint8_t               flags_end,
3597    uint8_t              *out_bytes
3598)
3599{
3600  __m128i h_vecs[8];
3601  __m128i counter_low_vec, counter_high_vec;
3602  uint8_t block_flags;
3603  size_t  block;
3604
3605  h_vecs[0] = set1_128_avx512(key[0]);
3606  h_vecs[1] = set1_128_avx512(key[1]);
3607  h_vecs[2] = set1_128_avx512(key[2]);
3608  h_vecs[3] = set1_128_avx512(key[3]);
3609  h_vecs[4] = set1_128_avx512(key[4]);
3610  h_vecs[5] = set1_128_avx512(key[5]);
3611  h_vecs[6] = set1_128_avx512(key[6]);
3612  h_vecs[7] = set1_128_avx512(key[7]);
3613
3614  load_counters4_avx512(counter, increment_counter, &counter_low_vec, &counter_high_vec);
3615  block_flags = flags | flags_start;
3616
3617  for (block = 0; block < blocks; block++) {
3618    __m128i block_len_vec;
3619    __m128i block_flags_vec;
3620    __m128i msg_vecs[16];
3621    __m128i v[16];
3622
3623    if (block + 1 == blocks) {
3624      block_flags |= flags_end;
3625    }
3626    block_len_vec   = set1_128_avx512(BLAKE3_BLOCK_LEN);
3627    block_flags_vec = set1_128_avx512(block_flags);
3628    transpose_msg_vecs4_avx512(inputs, block * BLAKE3_BLOCK_LEN, msg_vecs);
3629
3630    v[0]  = h_vecs[0];
3631    v[1]  = h_vecs[1];
3632    v[2]  = h_vecs[2];
3633    v[3]  = h_vecs[3];
3634    v[4]  = h_vecs[4];
3635    v[5]  = h_vecs[5];
3636    v[6]  = h_vecs[6];
3637    v[7]  = h_vecs[7];
3638    v[8]  = set1_128_avx512(IV[0]);
3639    v[9]  = set1_128_avx512(IV[1]);
3640    v[10] = set1_128_avx512(IV[2]);
3641    v[11] = set1_128_avx512(IV[3]);
3642    v[12] = counter_low_vec;
3643    v[13] = counter_high_vec;
3644    v[14] = block_len_vec;
3645    v[15] = block_flags_vec;
3646
3647    round_fn4_avx512(v, msg_vecs, 0);
3648    round_fn4_avx512(v, msg_vecs, 1);
3649    round_fn4_avx512(v, msg_vecs, 2);
3650    round_fn4_avx512(v, msg_vecs, 3);
3651    round_fn4_avx512(v, msg_vecs, 4);
3652    round_fn4_avx512(v, msg_vecs, 5);
3653    round_fn4_avx512(v, msg_vecs, 6);
3654
3655    h_vecs[0] = xor_128_avx512(v[0], v[8]);
3656    h_vecs[1] = xor_128_avx512(v[1], v[9]);
3657    h_vecs[2] = xor_128_avx512(v[2], v[10]);
3658    h_vecs[3] = xor_128_avx512(v[3], v[11]);
3659    h_vecs[4] = xor_128_avx512(v[4], v[12]);
3660    h_vecs[5] = xor_128_avx512(v[5], v[13]);
3661    h_vecs[6] = xor_128_avx512(v[6], v[14]);
3662    h_vecs[7] = xor_128_avx512(v[7], v[15]);
3663
3664    block_flags = flags;
3665  }
3666
3667  transpose_vecs_sse2(h_vecs);
3668  storeu_128_avx512(h_vecs[0], &out_bytes[0 * sizeof(__m128i)]);
3669  storeu_128_avx512(h_vecs[1], &out_bytes[2 * sizeof(__m128i)]);
3670  storeu_128_avx512(h_vecs[2], &out_bytes[4 * sizeof(__m128i)]);
3671  storeu_128_avx512(h_vecs[3], &out_bytes[6 * sizeof(__m128i)]);
3672  storeu_128_avx512(h_vecs[4], &out_bytes[1 * sizeof(__m128i)]);
3673  storeu_128_avx512(h_vecs[5], &out_bytes[3 * sizeof(__m128i)]);
3674  storeu_128_avx512(h_vecs[6], &out_bytes[5 * sizeof(__m128i)]);
3675  storeu_128_avx512(h_vecs[7], &out_bytes[7 * sizeof(__m128i)]);
3676}
3677
3678static inline void
3679hash_one_avx512(
3680    const uint8_t *input,
3681    size_t         blocks,
3682    const uint32_t key[8],
3683    uint64_t       counter,
3684    uint8_t        flags,
3685    uint8_t        flags_start,
3686    uint8_t        flags_end,
3687    uint8_t        out_bytes[BLAKE3_OUT_LEN]
3688)
3689{
3690  uint32_t cv[8];
3691  uint8_t  block_flags;
3692
3693  memcpy(cv, key, BLAKE3_KEY_LEN);
3694  block_flags = flags | flags_start;
3695  while (blocks > 0) {
3696    if (blocks == 1) {
3697      block_flags |= flags_end;
3698    }
3699    blake3_compress_in_place_avx512(cv, input, BLAKE3_BLOCK_LEN, counter, block_flags);
3700    input = &input[BLAKE3_BLOCK_LEN];
3701    blocks -= 1;
3702    block_flags = flags;
3703  }
3704  memcpy(out_bytes, cv, BLAKE3_OUT_LEN);
3705}
3706
3707void
3708blake3_hash_many_avx512(
3709    const uint8_t *const *inputs,
3710    size_t                num_inputs,
3711    size_t                blocks,
3712    const uint32_t        key[8],
3713    uint64_t              counter,
3714    int                   increment_counter,
3715    uint8_t               flags,
3716    uint8_t               flags_start,
3717    uint8_t               flags_end,
3718    uint8_t              *out_bytes
3719)
3720{
3721  while (num_inputs >= 16) {
3722    blake3_hash16_avx512(
3723        inputs, blocks, key, counter, increment_counter, flags, flags_start, flags_end, out_bytes
3724    );
3725    if (increment_counter) {
3726      counter += 16;
3727    }
3728    inputs += 16;
3729    num_inputs -= 16;
3730    out_bytes = &out_bytes[16 * BLAKE3_OUT_LEN];
3731  }
3732  while (num_inputs >= 8) {
3733    blake3_hash8_avx512(
3734        inputs, blocks, key, counter, increment_counter, flags, flags_start, flags_end, out_bytes
3735    );
3736    if (increment_counter) {
3737      counter += 8;
3738    }
3739    inputs += 8;
3740    num_inputs -= 8;
3741    out_bytes = &out_bytes[8 * BLAKE3_OUT_LEN];
3742  }
3743  while (num_inputs >= 4) {
3744    blake3_hash4_avx512(
3745        inputs, blocks, key, counter, increment_counter, flags, flags_start, flags_end, out_bytes
3746    );
3747    if (increment_counter) {
3748      counter += 4;
3749    }
3750    inputs += 4;
3751    num_inputs -= 4;
3752    out_bytes = &out_bytes[4 * BLAKE3_OUT_LEN];
3753  }
3754  while (num_inputs > 0) {
3755    hash_one_avx512(inputs[0], blocks, key, counter, flags, flags_start, flags_end, out_bytes);
3756    if (increment_counter) {
3757      counter += 1;
3758    }
3759    inputs += 1;
3760    num_inputs -= 1;
3761    out_bytes = &out_bytes[BLAKE3_OUT_LEN];
3762  }
3763}
3764#if defined(__clang__)
3765#pragma clang attribute pop
3766#elif defined(__GNUC__)
3767#pragma GCC pop_options
3768#endif
3769#endif
3770
3771/* dispatch functions */
3772void
3773blake3_compress_in_place(
3774    uint32_t      cv[8],
3775    const uint8_t block[BLAKE3_BLOCK_LEN],
3776    uint8_t       block_len,
3777    uint64_t      counter,
3778    uint8_t       flags
3779)
3780{
3781#if BLAKE3_X86_SIMD
3782  if (!blake3_cpu_detected)
3783    blake3_detect_cpu_features();
3784  if (blake3_cpu_features & AVX512) {
3785    blake3_compress_in_place_avx512(cv, block, block_len, counter, flags);
3786    return;
3787  }
3788  if (blake3_cpu_features & SSE41) {
3789    blake3_compress_in_place_sse41(cv, block, block_len, counter, flags);
3790    return;
3791  }
3792  if (blake3_cpu_features & SSE2) {
3793    blake3_compress_in_place_sse2(cv, block, block_len, counter, flags);
3794    return;
3795  }
3796#endif
3797  blake3_compress_in_place_portable(cv, block, block_len, counter, flags);
3798}
3799
3800void
3801blake3_compress_xof(
3802    const uint32_t cv[8],
3803    const uint8_t  block[BLAKE3_BLOCK_LEN],
3804    uint8_t        block_len,
3805    uint64_t       counter,
3806    uint8_t        flags,
3807    uint8_t        out[64]
3808)
3809{
3810#if BLAKE3_X86_SIMD
3811  if (!blake3_cpu_detected)
3812    blake3_detect_cpu_features();
3813  if (blake3_cpu_features & AVX512) {
3814    blake3_compress_xof_avx512(cv, block, block_len, counter, flags, out);
3815    return;
3816  }
3817  if (blake3_cpu_features & SSE41) {
3818    blake3_compress_xof_sse41(cv, block, block_len, counter, flags, out);
3819    return;
3820  }
3821  if (blake3_cpu_features & SSE2) {
3822    blake3_compress_xof_sse2(cv, block, block_len, counter, flags, out);
3823    return;
3824  }
3825#endif
3826  blake3_compress_xof_portable(cv, block, block_len, counter, flags, out);
3827}
3828
3829void
3830blake3_hash_many(
3831    const uint8_t *const *inputs,
3832    size_t                num_inputs,
3833    size_t                blocks,
3834    const uint32_t        key[8],
3835    uint64_t              counter,
3836    int                   increment_counter,
3837    uint8_t               flags,
3838    uint8_t               flags_start,
3839    uint8_t               flags_end,
3840    uint8_t              *out
3841)
3842{
3843#if BLAKE3_X86_SIMD
3844  if (!blake3_cpu_detected)
3845    blake3_detect_cpu_features();
3846  if (blake3_cpu_features & AVX512) {
3847    blake3_hash_many_avx512(
3848        inputs,
3849        num_inputs,
3850        blocks,
3851        key,
3852        counter,
3853        increment_counter,
3854        flags,
3855        flags_start,
3856        flags_end,
3857        out
3858    );
3859    return;
3860  }
3861  if (blake3_cpu_features & AVX2) {
3862    blake3_hash_many_avx2(
3863        inputs,
3864        num_inputs,
3865        blocks,
3866        key,
3867        counter,
3868        increment_counter,
3869        flags,
3870        flags_start,
3871        flags_end,
3872        out
3873    );
3874    return;
3875  }
3876  if (blake3_cpu_features & SSE41) {
3877    blake3_hash_many_sse41(
3878        inputs,
3879        num_inputs,
3880        blocks,
3881        key,
3882        counter,
3883        increment_counter,
3884        flags,
3885        flags_start,
3886        flags_end,
3887        out
3888    );
3889    return;
3890  }
3891  if (blake3_cpu_features & SSE2) {
3892    blake3_hash_many_sse2(
3893        inputs,
3894        num_inputs,
3895        blocks,
3896        key,
3897        counter,
3898        increment_counter,
3899        flags,
3900        flags_start,
3901        flags_end,
3902        out
3903    );
3904    return;
3905  }
3906#endif
3907  blake3_hash_many_portable(
3908      inputs,
3909      num_inputs,
3910      blocks,
3911      key,
3912      counter,
3913      increment_counter,
3914      flags,
3915      flags_start,
3916      flags_end,
3917      out
3918  );
3919}
3920
3921size_t
3922blake3_simd_degree(void)
3923{
3924#if BLAKE3_X86_SIMD
3925  if (!blake3_cpu_detected)
3926    blake3_detect_cpu_features();
3927  if (blake3_cpu_features & AVX512)
3928    return 16;
3929  if (blake3_cpu_features & AVX2)
3930    return 8;
3931  if (blake3_cpu_features & SSE41)
3932    return 4;
3933  if (blake3_cpu_features & SSE2)
3934    return 4;
3935#endif
3936  return 1;
3937}
3938
3939/* core hasher implementation */
3940const char *
3941blake3_version(void)
3942{
3943  return BLAKE3_VERSION_STRING;
3944}
3945
3946static inline void
3947chunk_state_init(struct Blake3ChunkState *self, const uint32_t key[8], uint8_t flags)
3948{
3949  memcpy(self->cv, key, BLAKE3_KEY_LEN);
3950  self->chunk_counter = 0;
3951  memset(self->buf, 0, BLAKE3_BLOCK_LEN);
3952  self->buf_len           = 0;
3953  self->blocks_compressed = 0;
3954  self->flags             = flags;
3955}
3956
3957static inline void
3958chunk_state_reset(struct Blake3ChunkState *self, const uint32_t key[8], uint64_t chunk_counter)
3959{
3960  memcpy(self->cv, key, BLAKE3_KEY_LEN);
3961  self->chunk_counter     = chunk_counter;
3962  self->blocks_compressed = 0;
3963  memset(self->buf, 0, BLAKE3_BLOCK_LEN);
3964  self->buf_len = 0;
3965}
3966
3967static inline size_t
3968chunk_state_len(const struct Blake3ChunkState *self)
3969{
3970  return (BLAKE3_BLOCK_LEN * (size_t)self->blocks_compressed) + ((size_t)self->buf_len);
3971}
3972
3973static inline size_t
3974chunk_state_fill_buf(struct Blake3ChunkState *self, const uint8_t *input, size_t input_len)
3975{
3976  size_t   take = BLAKE3_BLOCK_LEN - ((size_t)self->buf_len);
3977  uint8_t *dest;
3978
3979  if (take > input_len) {
3980    take = input_len;
3981  }
3982  dest = self->buf + ((size_t)self->buf_len);
3983  memcpy(dest, input, take);
3984  self->buf_len += (uint8_t)take;
3985  return take;
3986}
3987
3988static inline uint8_t
3989chunk_state_maybe_start_flag(const struct Blake3ChunkState *self)
3990{
3991  if (self->blocks_compressed == 0) {
3992    return CHUNK_START;
3993  } else {
3994    return 0;
3995  }
3996}
3997
3998static inline struct Output
3999make_output(
4000    const uint32_t input_cv[8],
4001    const uint8_t  block[BLAKE3_BLOCK_LEN],
4002    uint8_t        block_len,
4003    uint64_t       counter,
4004    uint8_t        flags
4005)
4006{
4007  struct Output ret;
4008
4009  memcpy(ret.input_cv, input_cv, 32);
4010  memcpy(ret.block, block, BLAKE3_BLOCK_LEN);
4011  ret.block_len = block_len;
4012  ret.counter   = counter;
4013  ret.flags     = flags;
4014  return ret;
4015}
4016
4017static inline void
4018output_chaining_value(const struct Output *self, uint8_t cv[32])
4019{
4020  uint32_t cv_words[8];
4021
4022  memcpy(cv_words, self->input_cv, 32);
4023  blake3_compress_in_place(cv_words, self->block, self->block_len, self->counter, self->flags);
4024  store_cv_words(cv, cv_words);
4025}
4026
4027static inline void
4028output_root_bytes(const struct Output *self, uint64_t seek, uint8_t *out, size_t out_len)
4029{
4030  uint64_t output_block_counter = seek / 64;
4031  size_t   offset_within_block  = seek % 64;
4032  uint8_t  wide_buf[64];
4033  size_t   available_bytes;
4034  size_t   memcpy_len;
4035
4036  while (out_len > 0) {
4037    blake3_compress_xof(
4038        self->input_cv,
4039        self->block,
4040        self->block_len,
4041        output_block_counter,
4042        self->flags | ROOT,
4043        wide_buf
4044    );
4045    available_bytes = 64 - offset_within_block;
4046    if (out_len > available_bytes) {
4047      memcpy_len = available_bytes;
4048    } else {
4049      memcpy_len = out_len;
4050    }
4051    memcpy(out, wide_buf + offset_within_block, memcpy_len);
4052    out += memcpy_len;
4053    out_len -= memcpy_len;
4054    output_block_counter += 1;
4055    offset_within_block = 0;
4056  }
4057}
4058
4059static inline void
4060chunk_state_update(struct Blake3ChunkState *self, const uint8_t *input, size_t input_len)
4061{
4062  size_t take;
4063
4064  if (self->buf_len > 0) {
4065    take = chunk_state_fill_buf(self, input, input_len);
4066    input += take;
4067    input_len -= take;
4068    if (input_len > 0) {
4069      blake3_compress_in_place(
4070          self->cv,
4071          self->buf,
4072          BLAKE3_BLOCK_LEN,
4073          self->chunk_counter,
4074          self->flags | chunk_state_maybe_start_flag(self)
4075      );
4076      self->blocks_compressed += 1;
4077      self->buf_len = 0;
4078      memset(self->buf, 0, BLAKE3_BLOCK_LEN);
4079    }
4080  }
4081
4082  while (input_len > BLAKE3_BLOCK_LEN) {
4083    blake3_compress_in_place(
4084        self->cv,
4085        input,
4086        BLAKE3_BLOCK_LEN,
4087        self->chunk_counter,
4088        self->flags | chunk_state_maybe_start_flag(self)
4089    );
4090    self->blocks_compressed += 1;
4091    input += BLAKE3_BLOCK_LEN;
4092    input_len -= BLAKE3_BLOCK_LEN;
4093  }
4094
4095  take = chunk_state_fill_buf(self, input, input_len);
4096  input += take;
4097  input_len -= take;
4098}
4099
4100static inline struct Output
4101chunk_state_output(const struct Blake3ChunkState *self)
4102{
4103  uint8_t block_flags = self->flags | chunk_state_maybe_start_flag(self) | CHUNK_END;
4104
4105  return make_output(self->cv, self->buf, self->buf_len, self->chunk_counter, block_flags);
4106}
4107
4108static inline struct Output
4109parent_output(const uint8_t block[BLAKE3_BLOCK_LEN], const uint32_t key[8], uint8_t flags)
4110{
4111  return make_output(key, block, BLAKE3_BLOCK_LEN, 0, flags | PARENT);
4112}
4113
4114static unsigned int
4115highest_one(uint64_t x)
4116{
4117#if defined(__GNUC__) || defined(__clang__)
4118  return 63 ^ __builtin_clzll(x);
4119#else
4120  unsigned int c = 0;
4121  if (x & 0xffffffff00000000ULL) {
4122    x >>= 32;
4123    c += 32;
4124  }
4125  if (x & 0x00000000ffff0000ULL) {
4126    x >>= 16;
4127    c += 16;
4128  }
4129  if (x & 0x000000000000ff00ULL) {
4130    x >>= 8;
4131    c += 8;
4132  }
4133  if (x & 0x00000000000000f0ULL) {
4134    x >>= 4;
4135    c += 4;
4136  }
4137  if (x & 0x000000000000000cULL) {
4138    x >>= 2;
4139    c += 2;
4140  }
4141  if (x & 0x0000000000000002ULL) {
4142    c += 1;
4143  }
4144  return c;
4145#endif
4146}
4147
4148static inline uint64_t
4149round_down_to_power_of_2(uint64_t x)
4150{
4151  return 1ULL << highest_one(x | 1);
4152}
4153
4154static inline size_t
4155left_len(size_t content_len)
4156{
4157  size_t full_chunks = (content_len - 1) / BLAKE3_CHUNK_LEN;
4158
4159  return round_down_to_power_of_2(full_chunks) * BLAKE3_CHUNK_LEN;
4160}
4161
4162static inline size_t
4163compress_chunks_parallel(
4164    const uint8_t *input,
4165    size_t         input_len,
4166    const uint32_t key[8],
4167    uint64_t       chunk_counter,
4168    uint8_t        flags,
4169    uint8_t       *out
4170)
4171{
4172  const uint8_t *chunks_array[MAX_SIMD_DEGREE];
4173  size_t         input_position   = 0;
4174  size_t         chunks_array_len = 0;
4175
4176  assert(0 < input_len);
4177  assert(input_len <= MAX_SIMD_DEGREE * BLAKE3_CHUNK_LEN);
4178
4179  while (input_len - input_position >= BLAKE3_CHUNK_LEN) {
4180    chunks_array[chunks_array_len] = &input[input_position];
4181    input_position += BLAKE3_CHUNK_LEN;
4182    chunks_array_len += 1;
4183  }
4184
4185  blake3_hash_many(
4186      chunks_array,
4187      chunks_array_len,
4188      BLAKE3_CHUNK_LEN / BLAKE3_BLOCK_LEN,
4189      key,
4190      chunk_counter,
4191      1,
4192      flags,
4193      CHUNK_START,
4194      CHUNK_END,
4195      out
4196  );
4197
4198  if (input_len > input_position) {
4199    uint64_t                counter = chunk_counter + (uint64_t)chunks_array_len;
4200    struct Blake3ChunkState chunk_state;
4201    struct Output           output;
4202
4203    chunk_state_init(&chunk_state, key, flags);
4204    chunk_state.chunk_counter = counter;
4205    chunk_state_update(&chunk_state, &input[input_position], input_len - input_position);
4206    output = chunk_state_output(&chunk_state);
4207    output_chaining_value(&output, &out[chunks_array_len * BLAKE3_OUT_LEN]);
4208    return chunks_array_len + 1;
4209  } else {
4210    return chunks_array_len;
4211  }
4212}
4213
4214static inline size_t
4215compress_parents_parallel(
4216    const uint8_t *child_chaining_values,
4217    size_t         num_chaining_values,
4218    const uint32_t key[8],
4219    uint8_t        flags,
4220    uint8_t       *out
4221)
4222{
4223  const uint8_t *parents_array[MAX_SIMD_DEGREE_OR_2];
4224  size_t         parents_array_len = 0;
4225
4226  assert(2 <= num_chaining_values);
4227  assert(num_chaining_values <= 2 * MAX_SIMD_DEGREE_OR_2);
4228
4229  while (num_chaining_values - (2 * parents_array_len) >= 2) {
4230    parents_array[parents_array_len] =
4231        &child_chaining_values[2 * parents_array_len * BLAKE3_OUT_LEN];
4232    parents_array_len += 1;
4233  }
4234
4235  blake3_hash_many(parents_array, parents_array_len, 1, key, 0, 0, flags | PARENT, 0, 0, out);
4236
4237  if (num_chaining_values > 2 * parents_array_len) {
4238    memcpy(
4239        &out[parents_array_len * BLAKE3_OUT_LEN],
4240        &child_chaining_values[2 * parents_array_len * BLAKE3_OUT_LEN],
4241        BLAKE3_OUT_LEN
4242    );
4243    return parents_array_len + 1;
4244  } else {
4245    return parents_array_len;
4246  }
4247}
4248
4249static inline size_t
4250blake3_compress_subtree_wide(
4251    const uint8_t *input,
4252    size_t         input_len,
4253    const uint32_t key[8],
4254    uint64_t       chunk_counter,
4255    uint8_t        flags,
4256    uint8_t       *out
4257)
4258{
4259  size_t degree;
4260
4261  if (input_len <= (size_t)blake3_simd_degree() * BLAKE3_CHUNK_LEN) {
4262    return compress_chunks_parallel(input, input_len, key, chunk_counter, flags, out);
4263  }
4264
4265  degree = blake3_simd_degree();
4266  if (degree > 1) {
4267    size_t child_len   = round_down_to_power_of_2(input_len - 1) / degree;
4268    size_t cvs_written = 0;
4269
4270    if (child_len < BLAKE3_CHUNK_LEN) {
4271      child_len = BLAKE3_CHUNK_LEN;
4272    }
4273
4274    while (input_len > 0) {
4275      size_t take = child_len;
4276      size_t sub_cvs;
4277
4278      if (take > input_len) {
4279        take = input_len;
4280      }
4281      sub_cvs = blake3_compress_subtree_wide(
4282          input, take, key, chunk_counter, flags, &out[cvs_written * BLAKE3_OUT_LEN]
4283      );
4284      cvs_written += sub_cvs;
4285      chunk_counter += take / BLAKE3_CHUNK_LEN;
4286      input += take;
4287      input_len -= take;
4288    }
4289
4290    while (cvs_written > 2) {
4291      cvs_written = compress_parents_parallel(out, cvs_written, key, flags, out);
4292    }
4293    return cvs_written;
4294  } else {
4295    /* fallback when simd degree is 1 */
4296    size_t left      = left_len(input_len);
4297    size_t right     = input_len - left;
4298    size_t left_cvs  = blake3_compress_subtree_wide(input, left, key, chunk_counter, flags, out);
4299    size_t right_cvs = blake3_compress_subtree_wide(
4300        input + left,
4301        right,
4302        key,
4303        chunk_counter + (left / BLAKE3_CHUNK_LEN),
4304        flags,
4305        &out[left_cvs * BLAKE3_OUT_LEN]
4306    );
4307
4308    return left_cvs + right_cvs;
4309  }
4310}
4311
4312static void
4313compress_subtree_to_parent_node(
4314    const uint8_t *input,
4315    size_t         input_len,
4316    const uint32_t key[8],
4317    uint64_t       chunk_counter,
4318    uint8_t        flags,
4319    uint8_t        out[2 * BLAKE3_OUT_LEN]
4320)
4321{
4322  uint8_t cv_array[2 * MAX_SIMD_DEGREE_OR_2 * BLAKE3_OUT_LEN];
4323  size_t  num_cvs =
4324      blake3_compress_subtree_wide(input, input_len, key, chunk_counter, flags, cv_array);
4325
4326  assert(num_cvs >= 2);
4327  while (num_cvs > 2) {
4328    num_cvs = compress_parents_parallel(cv_array, num_cvs, key, flags, cv_array);
4329  }
4330  memcpy(out, cv_array, 2 * BLAKE3_OUT_LEN);
4331}
4332
4333static inline void
4334hasher_init_base(struct Blake3Hasher *self, const uint32_t key[8], uint8_t flags)
4335{
4336  memcpy(self->key, key, BLAKE3_KEY_LEN);
4337  chunk_state_init(&self->chunk, key, flags);
4338  self->cv_stack_len = 0;
4339}
4340
4341void
4342blake3_hasher_init(struct Blake3Hasher *self)
4343{
4344  if (!blake3_cpu_detected)
4345    blake3_detect_cpu_features();
4346  hasher_init_base(self, IV, 0);
4347}
4348
4349static inline void
4350load_key_words(const uint8_t key[BLAKE3_KEY_LEN], uint32_t key_words[8])
4351{
4352  key_words[0] = load32(&key[0 * 4]);
4353  key_words[1] = load32(&key[1 * 4]);
4354  key_words[2] = load32(&key[2 * 4]);
4355  key_words[3] = load32(&key[3 * 4]);
4356  key_words[4] = load32(&key[4 * 4]);
4357  key_words[5] = load32(&key[5 * 4]);
4358  key_words[6] = load32(&key[6 * 4]);
4359  key_words[7] = load32(&key[7 * 4]);
4360}
4361
4362void
4363blake3_hasher_init_keyed(struct Blake3Hasher *self, const uint8_t key[BLAKE3_KEY_LEN])
4364{
4365  uint32_t key_words[8];
4366
4367  load_key_words(key, key_words);
4368  hasher_init_base(self, key_words, KEYED_HASH);
4369}
4370
4371void
4372blake3_hasher_init_derive_key_raw(
4373    struct Blake3Hasher *self, const void *context, size_t context_len
4374)
4375{
4376  struct Blake3Hasher context_hasher;
4377  uint8_t             context_key[BLAKE3_KEY_LEN];
4378  uint32_t            context_key_words[8];
4379
4380  hasher_init_base(&context_hasher, IV, DERIVE_KEY_CONTEXT);
4381  blake3_hasher_update(&context_hasher, context, context_len);
4382  blake3_hasher_finalize(&context_hasher, context_key, BLAKE3_KEY_LEN);
4383  load_key_words(context_key, context_key_words);
4384  hasher_init_base(self, context_key_words, DERIVE_KEY_MATERIAL);
4385}
4386
4387static inline unsigned int
4388popcnt(uint64_t x)
4389{
4390#if defined(__GNUC__) || defined(__clang__)
4391  return __builtin_popcountll(x);
4392#else
4393  unsigned int count = 0;
4394  while (x != 0) {
4395    count += 1;
4396    x &= x - 1;
4397  }
4398  return count;
4399#endif
4400}
4401
4402void
4403blake3_hasher_init_derive_key(struct Blake3Hasher *self, const char *context)
4404{
4405  blake3_hasher_init_derive_key_raw(self, context, strlen(context));
4406}
4407
4408static inline void
4409hasher_merge_cv_stack(struct Blake3Hasher *self, uint64_t total_len)
4410{
4411  size_t post_merge_stack_len = (size_t)popcnt(total_len);
4412
4413  while (self->cv_stack_len > post_merge_stack_len) {
4414    uint8_t      *parent_node = &self->cv_stack[(self->cv_stack_len - 2) * BLAKE3_OUT_LEN];
4415    struct Output output      = parent_output(parent_node, self->key, self->chunk.flags);
4416
4417    output_chaining_value(&output, parent_node);
4418    self->cv_stack_len -= 1;
4419  }
4420}
4421
4422static inline void
4423hasher_push_cv(struct Blake3Hasher *self, uint8_t new_cv[BLAKE3_OUT_LEN], uint64_t chunk_counter)
4424{
4425  hasher_merge_cv_stack(self, chunk_counter);
4426  memcpy(&self->cv_stack[self->cv_stack_len * BLAKE3_OUT_LEN], new_cv, BLAKE3_OUT_LEN);
4427  self->cv_stack_len += 1;
4428}
4429
4430void
4431blake3_hasher_update(struct Blake3Hasher *self, const void *input, size_t input_len)
4432{
4433  const uint8_t *input_bytes;
4434
4435  if (input_len == 0) {
4436    return;
4437  }
4438
4439  input_bytes = (const uint8_t *)input;
4440
4441  if (chunk_state_len(&self->chunk) > 0) {
4442    size_t take = BLAKE3_CHUNK_LEN - chunk_state_len(&self->chunk);
4443
4444    if (take > input_len) {
4445      take = input_len;
4446    }
4447    chunk_state_update(&self->chunk, input_bytes, take);
4448    input_bytes += take;
4449    input_len -= take;
4450    if (input_len > 0) {
4451      struct Output output = chunk_state_output(&self->chunk);
4452      uint8_t       chunk_cv[32];
4453
4454      output_chaining_value(&output, chunk_cv);
4455      hasher_push_cv(self, chunk_cv, self->chunk.chunk_counter);
4456      chunk_state_reset(&self->chunk, self->key, self->chunk.chunk_counter + 1);
4457    } else {
4458      return;
4459    }
4460  }
4461
4462  while (input_len > BLAKE3_CHUNK_LEN) {
4463    size_t   subtree_len  = round_down_to_power_of_2(input_len);
4464    uint64_t count_so_far = self->chunk.chunk_counter * BLAKE3_CHUNK_LEN;
4465    uint64_t subtree_chunks;
4466
4467    while ((((uint64_t)(subtree_len - 1)) & count_so_far) != 0) {
4468      subtree_len /= 2;
4469    }
4470    subtree_chunks = subtree_len / BLAKE3_CHUNK_LEN;
4471    if (subtree_len <= BLAKE3_CHUNK_LEN) {
4472      struct Blake3ChunkState chunk_state;
4473      struct Output           output;
4474      uint8_t                 cv[BLAKE3_OUT_LEN];
4475
4476      chunk_state_init(&chunk_state, self->key, self->chunk.flags);
4477      chunk_state.chunk_counter = self->chunk.chunk_counter;
4478      chunk_state_update(&chunk_state, input_bytes, subtree_len);
4479      output = chunk_state_output(&chunk_state);
4480      output_chaining_value(&output, cv);
4481      hasher_push_cv(self, cv, chunk_state.chunk_counter);
4482    } else {
4483      uint8_t cv_pair[2 * BLAKE3_OUT_LEN];
4484
4485      compress_subtree_to_parent_node(
4486          input_bytes, subtree_len, self->key, self->chunk.chunk_counter, self->chunk.flags, cv_pair
4487      );
4488      hasher_push_cv(self, cv_pair, self->chunk.chunk_counter);
4489      hasher_push_cv(
4490          self, &cv_pair[BLAKE3_OUT_LEN], self->chunk.chunk_counter + (subtree_chunks / 2)
4491      );
4492    }
4493    self->chunk.chunk_counter += subtree_chunks;
4494    input_bytes += subtree_len;
4495    input_len -= subtree_len;
4496  }
4497
4498  if (input_len > 0) {
4499    chunk_state_update(&self->chunk, input_bytes, input_len);
4500    hasher_merge_cv_stack(self, self->chunk.chunk_counter);
4501  }
4502}
4503
4504void
4505blake3_hasher_finalize_seek(
4506    const struct Blake3Hasher *self, uint64_t seek, uint8_t *out, size_t out_len
4507)
4508{
4509  struct Output output;
4510  size_t        cvs_remaining;
4511
4512  if (out_len == 0) {
4513    return;
4514  }
4515
4516  if (self->cv_stack_len == 0) {
4517    output = chunk_state_output(&self->chunk);
4518    output_root_bytes(&output, seek, out, out_len);
4519    return;
4520  }
4521
4522  if (chunk_state_len(&self->chunk) > 0) {
4523    cvs_remaining = self->cv_stack_len;
4524    output        = chunk_state_output(&self->chunk);
4525  } else {
4526    cvs_remaining = self->cv_stack_len - 2;
4527    output = parent_output(&self->cv_stack[cvs_remaining * 32], self->key, self->chunk.flags);
4528  }
4529
4530  while (cvs_remaining > 0) {
4531    uint8_t parent_block[BLAKE3_BLOCK_LEN];
4532
4533    cvs_remaining -= 1;
4534    memcpy(parent_block, &self->cv_stack[cvs_remaining * 32], 32);
4535    output_chaining_value(&output, &parent_block[32]);
4536    output = parent_output(parent_block, self->key, self->chunk.flags);
4537  }
4538  output_root_bytes(&output, seek, out, out_len);
4539}
4540
4541void
4542blake3_hasher_finalize(const struct Blake3Hasher *self, uint8_t *out, size_t out_len)
4543{
4544  blake3_hasher_finalize_seek(self, 0, out, out_len);
4545}
4546
4547void
4548blake3_hasher_reset(struct Blake3Hasher *self)
4549{
4550  chunk_state_reset(&self->chunk, self->key, 0);
4551  self->cv_stack_len = 0;
4552}
4553
4554/* utility implementation */
4555static unsigned char *out;
4556static size_t         outlen = BLAKE3_OUT_LEN;
4557
4558static void
4559usage(void)
4560{
4561  fprintf(stderr, "usage: %s [-bct] [-l length] [file ...]\n", argv0);
4562  exit(1);
4563}
4564
4565static int
4566sumfile(const char *name, FILE *file, unsigned char *out_buf, size_t out_len)
4567{
4568  char                buf[16384];
4569  struct Blake3Hasher ctx;
4570  size_t              len;
4571
4572  blake3_hasher_init(&ctx);
4573  do {
4574    len = fread(buf, 1, sizeof(buf), file);
4575    if (len > 0)
4576      blake3_hasher_update(&ctx, buf, len);
4577  } while (len == sizeof(buf));
4578
4579  if (ferror(file)) {
4580    fprintf(stderr, "%s: read %s: ", argv0, name);
4581    perror(NULL);
4582    return 1;
4583  }
4584  blake3_hasher_finalize(&ctx, out_buf, out_len);
4585  return 0;
4586}
4587
4588static int
4589sum(const char *name, FILE *file)
4590{
4591  size_t i;
4592
4593  if (sumfile(name, file, out, outlen) != 0)
4594    return 1;
4595  for (i = 0; i < outlen; i++)
4596    printf("%02x", out[i]);
4597  printf("  %s\n", name);
4598  return 0;
4599}
4600
4601static int
4602hexval(int c)
4603{
4604  if ('0' <= c && c <= '9')
4605    return c - '0';
4606  if ('a' <= c && c <= 'f')
4607    return c - 'a' + 10;
4608  if ('A' <= c && c <= 'F')
4609    return c - 'A' + 10;
4610  return -1;
4611}
4612
4613static int
4614checkfile(const char *name, const char *mode, const char *str, unsigned char *out_buf, size_t len)
4615{
4616  FILE  *file;
4617  int    c1, c2;
4618  size_t i;
4619
4620  file = fopen(name, mode);
4621  if (!file) {
4622    fprintf(stderr, "%s: open %s: ", argv0, name);
4623    perror(NULL);
4624    return 1;
4625  }
4626  sumfile(name, file, out_buf, len);
4627  fclose(file);
4628
4629  for (i = 0; i < len; i++) {
4630    c1 = hexval(str[i * 2]);
4631    c2 = hexval(str[i * 2 + 1]);
4632    if (c1 == -1 || c2 == -1) {
4633      fprintf(stderr, "%s: skipping invalid checksum line\n", argv0);
4634      return 1;
4635    }
4636    if (out_buf[i] != (c1 << 4 | c2)) {
4637      printf("%s: FAILED\n", name);
4638      return 1;
4639    }
4640  }
4641  printf("%s: OK\n", name);
4642  return 0;
4643}
4644
4645static int
4646check(const char *name, FILE *file)
4647{
4648  const char *mode;
4649  char        buf[8192], *pos, *end;
4650  size_t      len;
4651  int         ret = 0, skip = 0;
4652
4653  buf[sizeof(buf) - 2] = 0;
4654  while (fgets(buf, sizeof(buf), file)) {
4655    if (buf[sizeof(buf) - 2]) {
4656      fprintf(stderr, "%s: skipping line that is too long\n", argv0);
4657      buf[sizeof(buf) - 2] = 0;
4658      skip                 = 1;
4659      ret                  = 1;
4660      continue;
4661    }
4662    if (skip) {
4663      skip = 0;
4664      continue;
4665    }
4666    pos = strchr(buf, ' ');
4667    if (!pos || pos == buf || (pos[1] != ' ' && pos[1] != '*') || (pos - buf) & 1) {
4668      fprintf(stderr, "%s: skipping invalid checksum line\n", argv0);
4669      ret = 1;
4670      continue;
4671    }
4672    mode = pos[1] == ' ' ? "r" : "rb";
4673    len  = (pos - buf) / 2;
4674    if (len > outlen) {
4675      outlen = len;
4676      free(out);
4677      out = malloc(len);
4678      if (!out) {
4679        perror(argv0);
4680        return 1;
4681      }
4682    }
4683    *pos = '\0';
4684    pos += 2;
4685    end = strchr(pos, '\n');
4686    if (end)
4687      *end = '\0';
4688    ret |= checkfile(pos, mode, buf, out, len);
4689  }
4690  if (ferror(file)) {
4691    fprintf(stderr, "%s: read %s: ", argv0, name);
4692    perror(NULL);
4693    ret = 1;
4694  }
4695  return ret;
4696}
4697
4698// ?man b3sum: compute blake3 checksums
4699// ?man arguments: file ...
4700// ?man compute and check blake3 message digests
4701int
4702main(int argc, char *argv[])
4703{
4704  int (*func)(const char *, FILE *) = sum;
4705  FILE       *file;
4706  char       *end;
4707  const char *name, *mode = NULL;
4708  int         ret = 0;
4709
4710  ARGBEGIN
4711  {
4712    // ?man -b: read in binary mode
4713    case 'b':
4714      mode = "rb";
4715      break;
4716    // ?man -c: check blake3 sums from file
4717    case 'c':
4718      func = check;
4719      break;
4720    // ?man -l:str: -l length: output digest length in bytes
4721    case 'l':
4722      outlen = strtoul(EARGF(usage()), &end, 10);
4723      if (*end)
4724        usage();
4725      break;
4726    // ?man -t: read in text mode
4727    case 't':
4728      mode = "r";
4729      break;
4730    default:
4731      usage();
4732  }
4733  ARGEND
4734
4735  out = malloc(outlen);
4736  if (!out) {
4737    perror(NULL);
4738    return 1;
4739  }
4740
4741  if (argc == 0) {
4742    if (!mode || strcmp(mode, "r") == 0 || freopen(NULL, mode, stdin)) {
4743      ret |= func("<stdin>", stdin);
4744    } else {
4745      fprintf(stderr, "%s: reopen stdin: ", argv0);
4746      perror(NULL);
4747      ret = 1;
4748    }
4749  } else {
4750    if (!mode)
4751      mode = "r";
4752    for (; argc > 0; argc--, argv++) {
4753      name = *argv;
4754      file = fopen(name, mode);
4755      if (file) {
4756        ret |= func(name, file);
4757        fclose(file);
4758      } else {
4759        fprintf(stderr, "%s: open %s: ", argv0, name);
4760        perror(NULL);
4761        ret = 1;
4762      }
4763    }
4764  }
4765
4766  free(out);
4767  if (fshut(stdin, "<stdin>") | fshut(stdout, "<stdout>"))
4768    ret = 1;
4769
4770  return ret;
4771}