/*
 * Copyright (c) 2025 Google Inc. All rights reserved
 *
 * Permission is hereby granted, free of charge, to any person obtaining
 * a copy of this software and associated documentation files
 * (the "Software"), to deal in the Software without restriction,
 * including without limitation the rights to use, copy, modify, merge,
 * publish, distribute, sublicense, and/or sell copies of the Software,
 * and to permit persons to whom the Software is furnished to do so,
 * subject to the following conditions:
 *
 * The above copyright notice and this permission notice shall be
 * included in all copies or substantial portions of the Software.
 *
 * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
 * EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF
 * MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT.
 * IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY
 * CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT,
 * TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE
 * SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
 */

#define LOCAL_TRACE 0

#include <err.h>
#include <immintrin.h>
#include <kernel/thread.h>
#include <lk/macros.h>
#include <lk/trace.h>
#include <lk/types.h>
#include <platform.h>
#include <platform/random.h>
#include <string.h>

#define RAND_RETRIES 3000

#define US2NS(us) ((us) * (1000ULL))
#define MS2NS(ms) (US2NS(ms) * 1000ULL)

static status_t rand32_retry(int retries, uint32_t* rand) {
    int attempts = 0;

    bool printed_long_wait = false;
    lk_time_ns_t wait_start_time = current_time_ns();
    while (attempts <= retries) {
        if (_rdseed32_step(rand)) {
            return NO_ERROR;
        }

        if (!arch_ints_disabled()) {
            /*
             * Sleep when called from thread context. The kernel enables
             * interrupts about when it starts the scheduler, so checking
             * if interrupts are enabled is a good approximation for whether
             * it's safe to sleep.
             */
            thread_sleep_ns(X86_64_RDSEED_RNG_ENTROPY_SLEEP_NS);
        }

        if (!printed_long_wait) {
            lk_time_ns_t time_waited = current_time_ns() - wait_start_time;
            if (time_waited >= MS2NS(X86_64_RDSEED_RNG_LONG_WAIT_MS)) {
                dprintf(ALWAYS, "X86_64 DRNG waited for a long time: %llu ns\n",
                        time_waited);
                printed_long_wait = true;
            }
        }
        ++attempts;
    }

    return ERR_FAULT;
}

void platform_random_get_bytes(uint8_t* dest, size_t length) {
    lk_time_ns_t initial_start_time = current_time_ns();
    /* TODO(b/359346016) check if RDSEED is supported. */

    uint32_t* aligned_start = (uint32_t*)align((uintptr_t)dest, 4);

    uintptr_t prefix_len =
            (uintptr_t)((uintptr_t)aligned_start - (uintptr_t)dest);

    if (prefix_len > 0) {
        uint32_t temprand;
        if (rand32_retry(RAND_RETRIES, &temprand) != NO_ERROR) {
            dprintf(CRITICAL, "X86_64 DRNG failed\n");
            return;
        }

        memcpy(dest, &temprand, prefix_len);
    }

    uintptr_t words = ((uintptr_t)length - prefix_len) / 4;
    for (uintptr_t i = 0; i < words; ++i, ++aligned_start) {
        if (rand32_retry(RAND_RETRIES, aligned_start) != NO_ERROR) {
            dprintf(CRITICAL, "X86_64 DRNG failed\n");
            return;
        }
    }

    uintptr_t suffix_len = (uintptr_t)length - prefix_len - (words * 4);
    if (suffix_len > 0) {
        uint32_t temprand;
        if (rand32_retry(RAND_RETRIES, &temprand) != NO_ERROR) {
            dprintf(CRITICAL, "X86_64 DRNG failed\n");
            return;
        }

        uint8_t* suffix_start = (uint8_t*)((uintptr_t)aligned_start + words);
        memcpy(suffix_start, &temprand, suffix_len);
    }

    lk_time_ns_t total_time = current_time_ns() - initial_start_time;
    if (total_time >= MS2NS(X86_64_RDSEED_RNG_PRINT_MS)) {
        dprintf(INFO, "X86_64 DRNG total time for %zu bytes: %llu ns\n", length,
                total_time);
    }
}
