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

// Defining protected as public to create object of dng_image
#define protected public

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

#include "common.h"
#include "dng_ifd.h"
#include "dng_stream.h"
#include "dng_tag_codes.h"
#include "dng_tag_types.h"
#include "memutils.h"

uint8* startPtr = nullptr;
bool isTestInProgress = false;
const size_t page_size = getpagesize();
struct sigaction new_action, old_action;
int numPages = 0;

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

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

    // Check if SIGSEGV is coming from the expected offset.
    uintptr_t flaggedAddress = (uintptr_t)info->si_addr;
    bool isExpectedAddress = ((flaggedAddress - (numPages - 1) * page_size) |
                              (uintptr_t)startPtr) == ((uintptr_t)startPtr);
    startPtr = nullptr;
    if (isExpectedAddress) {
        exit(EXIT_VULNERABLE);
    }

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

uint32 maxTagCountForTTShort() {
    const dng_rect bounds(0 /* top */, 0 /* left */, 100 /* bottom */, 100 /* right */);
    dng_image img(bounds, 1 /* planes */, ttShort);
    return img.PixelRange();
}

int main() {
    // 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);

    // The vulnerable function 'ParseTag' is called with 'tagCount' and a 'stream' buffer where
    // 'stream' buffer is iterated. In each iteration, 2 bytes are read and iteration occurs for
    // 'tagCount' number of times. Hence '2 * tagCount' bytes are read in 'stream' buffer.
    // Further, an internal method 'CheckTagCount()' is called in vulnerable function 'ParseTag'
    // which returns false if 'tagCount' exceeds the maximum allowed value that is 'maxTagCount'.
    // For the current Poc,  'tagCount' in 'ParseTag' is passed as 'maxTagCount + 1' so that
    // 'CheckTagCount' returns false.
    // A 'stream' buffer is created. Bytes after '2 * maxTagCount' are made protected.
    // Without the fix, check on returns value of 'CheckTagCount' is missing, leading to read of
    // 'stream' buffer outside the allowed range that is '2 * maxTagCount' leading to SIGSEGV which
    // is detected and test fails.
    // With fix, check on 'CheckTagCount' is present and in 'ParseTag' returns false.
    uint32 maxTagCount = maxTagCountForTTShort();
    uint32 tagCount = maxTagCount + 1;
    if ((2 * tagCount) % page_size != 0) {
        return EXIT_FAILURE;
    }
    numPages = 2 * tagCount / page_size /* Compute the number of pages */ +
            1 /* Extra page (protected) */;
    startPtr = (uint8*)memalign(page_size, numPages * page_size);
    FAIL_CHECK(startPtr);
    uint8* crashPtr = startPtr + page_size * (numPages - 1);
    DISABLE_MEM_ACCESS(crashPtr, page_size); /* Protect the last page. */

    // Configure the params and '2 + startPtr' is passed as buffer pointer which will ensure that
    // only allowed bytes are unprotected.
    dng_stream stream(2 + startPtr, 2 * tagCount /* size */, 0 /* offset */);
    dng_ifd ifd;
    dng_host host;
    ifd.fSamplesPerPixel = maxTagCount;
    isTestInProgress = true;

    // The vulnerable function 'ParseTag' is called with a 'stream' buffer and a 'tagCount'.
    // Without the fix, due to a missing check on 'tagCount' for 'tagcode' (tcBitsPerSample,
    // tcExtraSample, and tcSampleFormat) in 'ParseTag', 'stream' buffer is read for protected
    // bytes.
    // The test will fail if check is missing in any of the 'tagCode' (tcBitsPerSample,
    // tcExtraSample, and tcSampleFormat).
    if (!(ifd.ParseTag(host, stream /* stream */, 0 /* parentCode */, tcBitsPerSample /* tagCode */,
                       ttShort /* tagType */, tagCount, 0 /* tagOffset*/)) &&
        !(ifd.ParseTag(host, stream /* stream */, 0 /* parentCode */, tcExtraSamples /* tagCode */,
                       ttShort /* tagType */, tagCount, 0 /* tagOffset*/)) &&
        !(ifd.ParseTag(host, stream /* stream */, 0 /* parentCode */, tcSampleFormat /* tagCode */,
                       ttShort /* tagType */, tagCount, 0 /* tagOffset*/))) {
        return EXIT_SUCCESS;
    };

    return EXIT_FAILURE;
}
