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