Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion doc/dev_ref/todo.rst
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,6 @@ Hardware Specific Optimizations
* GFNI implementations of ZFEC, others?
* NEON/VMX/LSX support for the SIMD based GHASH
* SIMD evaluation of SHA-2 and SHA-3 compression functions
* Improved Salsa implementations (SIMD_4x32, AVX2, AVX512, ...)
* Add CLMUL/PMULL implementations for CRC24
* Add support for ARMv8.4-A SHA-3 instructions
* Support POWER8 SHA-2 extensions (GH #1486 + #1487)
Expand Down
15 changes: 7 additions & 8 deletions src/lib/stream/chacha/chacha_avx2/chacha_avx2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,10 @@ void BOTAN_FN_ISA_AVX2 ChaCha::chacha_avx2_x8(uint8_t output[64 * 8], uint32_t s
SIMD_8x32::reset_registers();

BOTAN_ASSERT(rounds % 2 == 0, "Valid rounds");
const SIMD_8x32 CTR0 = SIMD_8x32(0, 1, 2, 3, 4, 5, 6, 7);

const uint32_t C = 0xFFFFFFFF - state[12];
// NOLINTNEXTLINE(*-implicit-bool-conversion)
const SIMD_8x32 CTR1 = SIMD_8x32(0, C < 1, C < 2, C < 3, C < 4, C < 5, C < 6, C < 7);
const SIMD_8x32 CTR_LO = SIMD_8x32::splat(state[12]) + SIMD_8x32(0, 1, 2, 3, 4, 5, 6, 7);
// Carry into the high counter word for lanes whose low word wrapped
const SIMD_8x32 CTR_HI = SIMD_8x32::splat(state[13]) - CTR_LO.unsigned_lt(SIMD_8x32::splat(state[12]));

SIMD_8x32 R00 = SIMD_8x32::splat(state[0]);
SIMD_8x32 R01 = SIMD_8x32::splat(state[1]);
Expand All @@ -34,8 +33,8 @@ void BOTAN_FN_ISA_AVX2 ChaCha::chacha_avx2_x8(uint8_t output[64 * 8], uint32_t s
SIMD_8x32 R09 = SIMD_8x32::splat(state[9]);
SIMD_8x32 R10 = SIMD_8x32::splat(state[10]);
SIMD_8x32 R11 = SIMD_8x32::splat(state[11]);
SIMD_8x32 R12 = SIMD_8x32::splat(state[12]) + CTR0;
SIMD_8x32 R13 = SIMD_8x32::splat(state[13]) + CTR1;
SIMD_8x32 R12 = CTR_LO;
SIMD_8x32 R13 = CTR_HI;
SIMD_8x32 R14 = SIMD_8x32::splat(state[14]);
SIMD_8x32 R15 = SIMD_8x32::splat(state[15]);

Expand Down Expand Up @@ -173,8 +172,8 @@ void BOTAN_FN_ISA_AVX2 ChaCha::chacha_avx2_x8(uint8_t output[64 * 8], uint32_t s
R09 += SIMD_8x32::splat(state[9]);
R10 += SIMD_8x32::splat(state[10]);
R11 += SIMD_8x32::splat(state[11]);
R12 += SIMD_8x32::splat(state[12]) + CTR0;
R13 += SIMD_8x32::splat(state[13]) + CTR1;
R12 += CTR_LO;
R13 += CTR_HI;
R14 += SIMD_8x32::splat(state[14]);
R15 += SIMD_8x32::splat(state[15]);

Expand Down
20 changes: 8 additions & 12 deletions src/lib/stream/chacha/chacha_avx512/chacha_avx512.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,15 +14,11 @@ namespace Botan {
//static
void BOTAN_FN_ISA_AVX512 ChaCha::chacha_avx512_x16(uint8_t output[64 * 16], uint32_t state[16], size_t rounds) {
BOTAN_ASSERT(rounds % 2 == 0, "Valid rounds");
const SIMD_16x32 CTR0 = SIMD_16x32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);

const uint32_t C = 0xFFFFFFFF - state[12];

// clang-format off
const SIMD_16x32 CTR1 = SIMD_16x32(
// NOLINTNEXTLINE(*-implicit-bool-conversion)
0, C < 1, C < 2, C < 3, C < 4, C < 5, C < 6, C < 7, C < 8, C < 9, C < 10, C < 11, C < 12, C < 13, C < 14, C < 15);
// clang-format on
const SIMD_16x32 CTR_LO =
SIMD_16x32::splat(state[12]) + SIMD_16x32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);
// Carry into the high counter word for lanes whose low word wrapped
const SIMD_16x32 CTR_HI = SIMD_16x32::splat(state[13]) - CTR_LO.unsigned_lt(SIMD_16x32::splat(state[12]));

SIMD_16x32 R00 = SIMD_16x32::splat(state[0]);
SIMD_16x32 R01 = SIMD_16x32::splat(state[1]);
Expand All @@ -36,8 +32,8 @@ void BOTAN_FN_ISA_AVX512 ChaCha::chacha_avx512_x16(uint8_t output[64 * 16], uint
SIMD_16x32 R09 = SIMD_16x32::splat(state[9]);
SIMD_16x32 R10 = SIMD_16x32::splat(state[10]);
SIMD_16x32 R11 = SIMD_16x32::splat(state[11]);
SIMD_16x32 R12 = SIMD_16x32::splat(state[12]) + CTR0;
SIMD_16x32 R13 = SIMD_16x32::splat(state[13]) + CTR1;
SIMD_16x32 R12 = CTR_LO;
SIMD_16x32 R13 = CTR_HI;
SIMD_16x32 R14 = SIMD_16x32::splat(state[14]);
SIMD_16x32 R15 = SIMD_16x32::splat(state[15]);

Expand Down Expand Up @@ -175,8 +171,8 @@ void BOTAN_FN_ISA_AVX512 ChaCha::chacha_avx512_x16(uint8_t output[64 * 16], uint
R09 += SIMD_16x32::splat(state[9]);
R10 += SIMD_16x32::splat(state[10]);
R11 += SIMD_16x32::splat(state[11]);
R12 += SIMD_16x32::splat(state[12]) + CTR0;
R13 += SIMD_16x32::splat(state[13]) + CTR1;
R12 += CTR_LO;
R13 += CTR_HI;
R14 += SIMD_16x32::splat(state[14]);
R15 += SIMD_16x32::splat(state[15]);

Expand Down
16 changes: 7 additions & 9 deletions src/lib/stream/chacha/chacha_simd32/chacha_simd32.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,12 +14,10 @@ namespace Botan {
//static
void BOTAN_FN_ISA_SIMD_4X32 ChaCha::chacha_simd32_x4(uint8_t output[64 * 4], uint32_t state[16], size_t rounds) {
BOTAN_ASSERT(rounds % 2 == 0, "Valid rounds");
const SIMD_4x32 CTR0 = SIMD_4x32(0, 1, 2, 3);

const uint32_t C = 0xFFFFFFFF - state[12];

// NOLINTNEXTLINE(*-implicit-bool-conversion)
const SIMD_4x32 CTR1 = SIMD_4x32(0, C < 1, C < 2, C < 3);
const SIMD_4x32 CTR_LO = SIMD_4x32::splat(state[12]) + SIMD_4x32(0, 1, 2, 3);
// Carry into the high counter word for lanes whose low word wrapped
const SIMD_4x32 CTR_HI = SIMD_4x32::splat(state[13]) - CTR_LO.unsigned_lt(SIMD_4x32::splat(state[12]));

SIMD_4x32 R00 = SIMD_4x32::splat(state[0]);
SIMD_4x32 R01 = SIMD_4x32::splat(state[1]);
Expand All @@ -33,8 +31,8 @@ void BOTAN_FN_ISA_SIMD_4X32 ChaCha::chacha_simd32_x4(uint8_t output[64 * 4], uin
SIMD_4x32 R09 = SIMD_4x32::splat(state[9]);
SIMD_4x32 R10 = SIMD_4x32::splat(state[10]);
SIMD_4x32 R11 = SIMD_4x32::splat(state[11]);
SIMD_4x32 R12 = SIMD_4x32::splat(state[12]) + CTR0;
SIMD_4x32 R13 = SIMD_4x32::splat(state[13]) + CTR1;
SIMD_4x32 R12 = CTR_LO;
SIMD_4x32 R13 = CTR_HI;
SIMD_4x32 R14 = SIMD_4x32::splat(state[14]);
SIMD_4x32 R15 = SIMD_4x32::splat(state[15]);

Expand Down Expand Up @@ -172,8 +170,8 @@ void BOTAN_FN_ISA_SIMD_4X32 ChaCha::chacha_simd32_x4(uint8_t output[64 * 4], uin
R09 += SIMD_4x32::splat(state[9]);
R10 += SIMD_4x32::splat(state[10]);
R11 += SIMD_4x32::splat(state[11]);
R12 += SIMD_4x32::splat(state[12]) + CTR0;
R13 += SIMD_4x32::splat(state[13]) + CTR1;
R12 += CTR_LO;
R13 += CTR_HI;
R14 += SIMD_4x32::splat(state[14]);
R15 += SIMD_4x32::splat(state[15]);

Expand Down
132 changes: 113 additions & 19 deletions src/lib/stream/salsa20/salsa20.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,10 @@
#include <botan/internal/loadstor.h>
#include <botan/internal/rotate.h>

#if defined(BOTAN_HAS_CPUID)
#include <botan/internal/cpuid.h>
#endif

namespace Botan {

namespace {
Expand Down Expand Up @@ -122,6 +126,88 @@ void Salsa20::salsa_core(uint8_t output[64], const uint32_t input[16], size_t ro
store_le(x15 + input[15], output + 4 * 15);
}

size_t Salsa20::parallelism() {
#if defined(BOTAN_HAS_SALSA20_AVX512)
if(CPUID::has(CPUID::Feature::AVX512)) {
return 16;
}
#endif

#if defined(BOTAN_HAS_SALSA20_AVX2)
if(CPUID::has(CPUID::Feature::AVX2)) {
return 8;
}
#endif

return 4;
}

std::string Salsa20::provider() const {
#if defined(BOTAN_HAS_SALSA20_AVX512)
if(auto feat = CPUID::check(CPUID::Feature::AVX512)) {
return *feat;
}
#endif

#if defined(BOTAN_HAS_SALSA20_AVX2)
if(auto feat = CPUID::check(CPUID::Feature::AVX2)) {
return *feat;
}
#endif

#if defined(BOTAN_HAS_SALSA20_SIMD32)
if(auto feat = CPUID::check(CPUID::Feature::SIMD_4X32)) {
return *feat;
}
#endif

return "base";
}

//static
void Salsa20::salsa20(uint8_t output[], size_t output_blocks, uint32_t state[16], size_t rounds) {
BOTAN_ASSERT(rounds % 2 == 0, "Valid rounds");

#if defined(BOTAN_HAS_SALSA20_AVX512)
if(CPUID::has(CPUID::Feature::AVX512)) {
while(output_blocks >= 16) {
Salsa20::salsa20_avx512_x16(output, state, rounds);
output += 16 * 64;
output_blocks -= 16;
}
}
#endif

#if defined(BOTAN_HAS_SALSA20_AVX2)
if(CPUID::has(CPUID::Feature::AVX2)) {
while(output_blocks >= 8) {
Salsa20::salsa20_avx2_x8(output, state, rounds);
output += 8 * 64;
output_blocks -= 8;
}
}
#endif

#if defined(BOTAN_HAS_SALSA20_SIMD32)
if(CPUID::has(CPUID::Feature::SIMD_4X32)) {
while(output_blocks >= 4) {
Salsa20::salsa20_simd32_x4(output, state, rounds);
output += 4 * 64;
output_blocks -= 4;
}
}
#endif

for(size_t i = 0; i != output_blocks; ++i) {
salsa_core(output + 64 * i, state, rounds);

++state[8];
if(state[8] == 0) {
state[9] += 1;
}
}
}

/*
* Combine cipher stream with message
*/
Expand All @@ -132,12 +218,7 @@ void Salsa20::cipher_bytes(const uint8_t in[], uint8_t out[], size_t length) {
const size_t available = m_buffer.size() - m_position;

xor_buf(out, in, &m_buffer[m_position], available);
salsa_core(m_buffer.data(), m_state.data(), 20);

++m_state[8];
if(m_state[8] == 0) {
m_state[9] += 1;
}
salsa20(m_buffer.data(), m_buffer.size() / 64, m_state.data(), 20);

length -= available;
in += available;
Expand All @@ -151,6 +232,27 @@ void Salsa20::cipher_bytes(const uint8_t in[], uint8_t out[], size_t length) {
m_position += length;
}

void Salsa20::generate_keystream(uint8_t out[], size_t length) {
assert_key_material_set();

while(length >= m_buffer.size() - m_position) {
const size_t available = m_buffer.size() - m_position;

// TODO: this could write directly to the output buffer
// instead of bouncing it through m_buffer first
copy_mem(out, &m_buffer[m_position], available);
salsa20(m_buffer.data(), m_buffer.size() / 64, m_state.data(), 20);

length -= available;
out += available;
m_position = 0;
}

copy_mem(out, &m_buffer[m_position], length);

m_position += length;
}

void Salsa20::initialize_state() {
static const uint32_t TAU[] = {0x61707865, 0x3120646e, 0x79622d36, 0x6b206574};

Expand Down Expand Up @@ -205,7 +307,9 @@ void Salsa20::key_schedule(std::span<const uint8_t> key) {
load_le<uint32_t>(m_key.data(), key.data(), m_key.size());

m_state.resize(16);
m_buffer.resize(64);

const size_t salsa_block = 64;
m_buffer.resize(parallelism() * salsa_block);

set_iv(nullptr, 0);
}
Expand Down Expand Up @@ -255,12 +359,7 @@ void Salsa20::set_iv_bytes(const uint8_t iv[], size_t length) {
m_state[8] = 0;
m_state[9] = 0;

salsa_core(m_buffer.data(), m_state.data(), 20);
++m_state[8];
if(m_state[8] == 0) {
m_state[9] += 1;
}

salsa20(m_buffer.data(), m_buffer.size() / 64, m_state.data(), 20);
m_position = 0;
}

Expand Down Expand Up @@ -302,12 +401,7 @@ void Salsa20::seek(uint64_t offset) {
m_state[8] = static_cast<uint32_t>(counter);
m_state[9] = static_cast<uint32_t>(counter >> 32);

salsa_core(m_buffer.data(), m_state.data(), 20);

++m_state[8];
if(m_state[8] == 0) {
m_state[9] += 1;
}
salsa20(m_buffer.data(), m_buffer.size() / 64, m_state.data(), 20);

m_position = offset % 64;
}
Expand Down
18 changes: 18 additions & 0 deletions src/lib/stream/salsa20/salsa20.h
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ namespace Botan {
*/
class Salsa20 final : public StreamCipher {
public:
std::string provider() const override;
bool valid_iv_length(size_t iv_len) const override;
size_t default_iv_length() const override;
Key_Length_Specification key_spec() const override;
Expand All @@ -41,13 +42,30 @@ class Salsa20 final : public StreamCipher {

protected:
void cipher_bytes(const uint8_t in[], uint8_t out[], size_t length) override;
void generate_keystream(uint8_t out[], size_t len) override;
void set_iv_bytes(const uint8_t iv[], size_t iv_len) override;

private:
void key_schedule(std::span<const uint8_t> key) override;

void initialize_state();

static size_t parallelism();

static void salsa20(uint8_t output[], size_t output_blocks, uint32_t state[16], size_t rounds);

#if defined(BOTAN_HAS_SALSA20_SIMD32)
static void salsa20_simd32_x4(uint8_t output[64 * 4], uint32_t state[16], size_t rounds);
#endif

#if defined(BOTAN_HAS_SALSA20_AVX2)
static void salsa20_avx2_x8(uint8_t output[64 * 8], uint32_t state[16], size_t rounds);
#endif

#if defined(BOTAN_HAS_SALSA20_AVX512)
static void salsa20_avx512_x16(uint8_t output[64 * 16], uint32_t state[16], size_t rounds);
#endif

secure_vector<uint32_t> m_key;
secure_vector<uint32_t> m_state;
secure_vector<uint8_t> m_buffer;
Expand Down
17 changes: 17 additions & 0 deletions src/lib/stream/salsa20/salsa20_avx2/info.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
<internal_defines>
SALSA20_AVX2 -> 20260726
</internal_defines>

<module_info>
name -> "Salsa20 AVX2"
brief -> "Salsa20 using AVX2 instructions"
</module_info>

<isa>
avx2
</isa>

<requires>
simd_avx2
cpuid
</requires>
Loading
Loading