/*
 * Copyright (C) 2025 The Android Open Source Project
 *
 * Licensed under the Apache License, Version 2.0 (the "License");
 * you may not use this file except in compliance with the License.
 * You may obtain a copy of the License at
 *
 *      http://www.apache.org/licenses/LICENSE-2.0
 *
 * Unless required by applicable law or agreed to in writing, software
 * distributed under the License is distributed on an "AS IS" BASIS,
 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
 * See the License for the specific language governing permissions and
 * limitations under the License.
 */

#include <sys/mman.h>
#include <unistd.h>
#include <utils/Errors.h>

#include "binary_loader.h"
#include "common.h"
#include "dng_exceptions.h"
#include "memutils.h"

typedef int (*HuffDecode_func)(void*, void*);
struct sigaction new_action, old_action;
uintptr_t faultPtr = (uintptr_t)nullptr;
const int dngDecoderSize = 300 /* Arbitrary size of buffer */;
const int pageSize = getpagesize();
bool isTestInProgress = false;
char* startPtr = nullptr;
void* dngDecoder = nullptr;
HuffDecode_func HuffDecode_ptr;

void releaseMemoryAndExit(android::status_t status) {
    free(dngDecoder);
    free(startPtr);
    dngDecoder = nullptr;
    startPtr = nullptr;
    exit(status);
}

void sig_handler(int signum, siginfo_t* info, void* context) {
    if (info->si_addr == nullptr) {
        (*old_action.sa_sigaction)(signum, info, context);
        return;
    }

    // Ignore the signals that occur when the test has not started.
    if (!isTestInProgress) {
        (*old_action.sa_sigaction)(signum, info, context);
        return;
    }

    // When 'HuffDecode()' is called for the second time with 'nullptr', NPD occurs with the offset
    // equals to value of 'faultPtr' and test fails.
    if ((uintptr_t)info->si_addr == faultPtr) {
        releaseMemoryAndExit(EXIT_VULNERABLE);
    }

    // When 'HuffDecode()' is called for the first time with guarded buffer, 'SIGSEGV' occurs when
    // buffer is read. Compute the offset of the byte read from the 'startPtr' address.
    uintptr_t offset = ((uintptr_t)info->si_addr ^ (uintptr_t)startPtr) & (uintptr_t)info->si_addr;

    // Check if signal is coming from guarded buffer and store the offset in 'faultPtr'.
    if (offset < pageSize) {
        faultPtr = offset;

        // Free the guarded buffer.
        ENABLE_MEM_ACCESS(startPtr, pageSize);
        free(startPtr);
        startPtr = nullptr;

        // Ignore the first call signal.
        return;
    }

    // Assumption failure for unexpected signals.
    releaseMemoryAndExit(EXIT_FAILURE);
}

int main(int /* argc */, char* argv[]) {
    // Setup signal handler.
    sigemptyset(&new_action.sa_mask);
    sigaddset(&new_action.sa_mask, SIGSEGV); // Add SIGSEGV to the signal mask
    sigaddset(&new_action.sa_mask, SIGBUS);  // Add SIGBUS to the signal mask
    new_action.sa_flags = SA_SIGINFO;
    new_action.sa_sigaction = sig_handler;
    sigaction(SIGSEGV, &new_action, &old_action);
    sigaction(SIGBUS, &new_action, &old_action);

    // Get the path to the shared library and offset from command-line arguments
    const char* libPath = argv[1];
    const uintptr_t functionOffset = strtoul(argv[2], nullptr, 0);

    // Get function address of 'HuffDecode()' from loaded library
    BinaryLoader binaryLoader(libPath);
    const uintptr_t functionAddress = binaryLoader.getFunctionAddress(functionOffset);
    FAIL_CHECK(functionAddress);

    // Create function pointer to 'HuffDecode()'
    HuffDecode_ptr = (HuffDecode_func)functionAddress;

    // Create a guarded buffer.
    startPtr = (char*)memalign(pageSize, pageSize);
    FAIL_CHECK(startPtr);
    memset(startPtr, 0, pageSize);
    DISABLE_MEM_ACCESS(startPtr, pageSize);

    // Call the vulnerable function 'HuffDecode()' for the first time with a guarded buffer.
    // 'SIGSEGV' is raised when the guarded buffer 'startPtr' is read in the vulnerable function.
    // 'faultPtr' stores the offset of the byte read from the guarded buffer 'startPtr' address.
    dngDecoder = malloc(dngDecoderSize);
    FAIL_CHECK(dngDecoder);
    memset(dngDecoder, 8 /* Set bitsLeft */, dngDecoderSize /* Size of Buffer */);
    isTestInProgress = true;
    HuffDecode_ptr(dngDecoder, startPtr);

    // Call the vulnerable function 'HuffDecode()' for second time with the 'nullptr'.
    // Without the fix, due to missing null pointer check, NPD occurred with the offset
    // equals to value of 'faultPtr'. With fix, 'dng_exception' is thrown.
    try {
        memset(dngDecoder, 8 /* Set bitsLeft */, dngDecoderSize /* Size of Buffer */);
        HuffDecode_ptr(dngDecoder, nullptr);
    } catch (const dng_exception& e) {
        // Dng exception is thrown in case of fix.
        releaseMemoryAndExit(EXIT_SUCCESS);
    }

    // Free the allocated memory.
    releaseMemoryAndExit(EXIT_FAILURE);
}
