/*
 * 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.
 */

// Define private as public to create SkEncodedInfo.
#define private public

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

#include "binary_loader.h"
#include "common.h"
#include "include/core/SkStream.h"
#include "memutils.h"
#include "src/codec/SkBmpStandardCodec.h"

bool isTestInProgress = false;
char* startPtr = nullptr;
char* crashPtr = nullptr;
struct sigaction new_action, old_action;
typedef int (*onPrepareToDecode_func)(SkBmpStandardCodec*, SkImageInfo const&,
                                      SkCodec::Options const&);
typedef int (*decodeRows_func)(SkBmpStandardCodec*, SkImageInfo const&, void*, unsigned long,
                               SkCodec::Options const&);

void sig_handler(int signum, siginfo_t* info, void* context) {
    // Free the allocated buffer.
    if (startPtr) {
        ENABLE_MEM_ACCESS(startPtr, getpagesize());
        free(startPtr);
        startPtr = nullptr;
    }

    // Check if test is in progress.
    if (!isTestInProgress) {
        (*old_action.sa_sigaction)(signum, info, context);
        return;
    }

    // Check if SIGSEGV is coming from expected address.
    if ((((uintptr_t)info->si_addr ^ (uintptr_t)crashPtr) & (uintptr_t)info->si_addr) == 0) {
        exit(EXIT_VULNERABLE);
    }

    // Assumption failure for unexpected signals.
    exit(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 offsets from command-line arguments
    const char* libPath = argv[1];
    const uintptr_t onPrepareToDecodeOffset = strtoul(argv[2], nullptr, 0);
    const uintptr_t decodeRowsOffset = strtoul(argv[3], nullptr, 0);

    // Get function address of 'onPrepareToDecode()' and 'decodeRows' from loaded library.
    BinaryLoader binaryLoader(libPath);
    const uintptr_t onPrepareToDecodeAddress =
            binaryLoader.getFunctionAddress(onPrepareToDecodeOffset);
    FAIL_CHECK(onPrepareToDecodeAddress);
    const uintptr_t decodeRowsAddress = binaryLoader.getFunctionAddress(decodeRowsOffset);
    FAIL_CHECK(decodeRowsAddress);

    // Create function pointer of 'onPrepareToDecode()' and 'decodeRows'.
    onPrepareToDecode_func onPrepareToDecode_ptr = (onPrepareToDecode_func)onPrepareToDecodeAddress;
    decodeRows_func decodeRows_ptr = (decodeRows_func)decodeRowsAddress;

    // Create guarded buffer.
    const int imageBufferSize = 2048;
    const size_t page_size = getpagesize();
    startPtr = (char*)memalign(page_size, 2 * page_size);
    FAIL_CHECK(startPtr);
    crashPtr = startPtr + page_size;
    DISABLE_MEM_ACCESS(crashPtr, page_size); /* Protect the last page. */
    char* bufferPtr = (startPtr + page_size) /* End of first page */ - imageBufferSize;

    // Configure the params of 'encodedInfo' and create stream buffer.
    SkEncodedInfo encodedInfo(32 /* width */, 32 /* height */,
                              SkEncodedInfo::kPalette_Color /* color */,
                              SkEncodedInfo::kUnpremul_Alpha /* Alpha Type */,
                              4 /* bitsPerComponent */, 8 /* colorDepth */,
                              nullptr /* ICC profile */);
    auto stream = std::make_unique<SkMemoryStream>(bufferPtr, imageBufferSize, false);

    // Configure the params of 'skBmpStandardCodec'.
    SkBmpStandardCodec skBmpStandardCodec(std::move(encodedInfo) /* SkEncodedInfo */,
                                          std::move(stream) /* SkStream */, 8 /* bitsPerPixel */,
                                          0 /* numColors */, 4 /* bytesPerColor */,
                                          1024 /* offset */,
                                          SkCodec::kTopDown_SkScanlineOrder /* rowOrder */,
                                          true /* isOpaque */, false /* inIco */);
    skBmpStandardCodec.fXformTime = SkBmpBaseCodec::kPalette_XformTime;

    // Configure the 'SkCodec::Options' and 'SkImageInfo'.
    SkCodec::Options opt;
    SkImageInfo imageInfo =
            SkImageInfo::Make(32 /* width */, 32 /* height */, kRGB_565_SkColorType /* color */,
                              kPremul_SkAlphaType /* alpha type */
            );

    // Call the function 'onPrepareToDecode' which further calls the vulnerable function
    // 'initializeSwizzler'. Without the fix, the vulnerable function 'initializeSwizzler' has
    // check on 'colorXform' instead of 'xformOnDecode' which causes 'fSwizzler' to hold function
    // address of 'swizzle_small_index_to_n32'. 'decodeRows' is then called, which calls
    // 'swizzle_small_index_to_n32' where OOB write on 'bufferPtr' occurs and test fails.
    onPrepareToDecode_ptr(&skBmpStandardCodec /* object */, imageInfo /* dstInfo */,
                          opt /* opts */);
    isTestInProgress = true;
    decodeRows_ptr(&skBmpStandardCodec /* object */, imageInfo /* dstInfo */,
                   (void*)bufferPtr /* dst  */, 64 /* dstRowBytes */, opt /* opts */);
    exit(EXIT_SUCCESS);
}
