#include <stdlib.h>
#include <string.h>
#include <stdio.h>
#include <stdint.h>
#include "../../sphincsplus/ref/randombytes.h"
#include "../../sphincsplus/ref/api.h"
#include "../../sphincsplus/ref/fors.h"
#include "../../sphincsplus/ref/hash.h"
#include "../../sphincsplus/ref/thash.h"
#include "../../sphincsplus/ref/utils.h"
#include "../../sphincsplus/ref/address.h"
#include "libbitcoinpqc/slh_dsa.h"
static const uint8_t *g_random_data = NULL;
static size_t g_random_data_size = 0;
static size_t g_random_data_offset = 0;
void slh_dsa_init_random_source(const uint8_t *random_data, size_t random_data_size) {
g_random_data = random_data;
g_random_data_size = random_data_size;
g_random_data_offset = 0;
}
void slh_dsa_setup_custom_random() {
}
void slh_dsa_restore_original_random() {
g_random_data = NULL;
g_random_data_size = 0;
g_random_data_offset = 0;
}
void custom_slh_randombytes_impl(uint8_t *out, size_t outlen) {
if (g_random_data == NULL || g_random_data_size == 0) {
FILE *f = fopen("/dev/urandom", "r");
if (!f) {
memset(out, 0, outlen);
return;
}
if (fread(out, 1, outlen, f) != outlen) {
memset(out, 0, outlen);
}
fclose(f);
return;
}
size_t remaining = g_random_data_size - g_random_data_offset;
if (outlen > remaining) {
size_t position = 0;
while (position < outlen) {
size_t to_copy = (outlen - position < remaining) ? outlen - position : remaining;
memcpy(out + position, g_random_data + g_random_data_offset, to_copy);
position += to_copy;
g_random_data_offset = (g_random_data_offset + to_copy) % g_random_data_size;
remaining = g_random_data_size - g_random_data_offset;
}
} else {
memcpy(out, g_random_data + g_random_data_offset, outlen);
g_random_data_offset = (g_random_data_offset + outlen) % g_random_data_size;
}
}
void slh_dsa_derandomize(uint8_t *seed, const uint8_t *m, size_t mlen, const uint8_t *sk) {
size_t combined_len = mlen + CRYPTO_SECRETKEYBYTES;
uint8_t *combined = malloc(combined_len);
if (combined) {
memcpy(combined, sk, CRYPTO_SECRETKEYBYTES);
memcpy(combined + CRYPTO_SECRETKEYBYTES, m, mlen);
uint8_t buffer[64] = {0};
for (size_t i = 0; i < combined_len; i++) {
buffer[i % 64] ^= combined[i];
}
for (size_t i = 0; i < 10; i++) {
for (size_t j = 0; j < 64; j++) {
buffer[j] = buffer[(j + 1) % 64] ^ buffer[(j + 7) % 64] ^ buffer[(j + 13) % 64];
}
}
memcpy(seed, buffer, 64);
memset(combined, 0, combined_len);
free(combined);
} else {
memset(seed, 0, 64);
}
}